extern crate alloc;
use alloc::collections::BTreeMap;
use alloc::collections::btree_map::Entry;
use rat_rdp_pdu::codecs::rfx::EntropyAlgorithm;
use rat_rdp_pdu::codecs::rfx::progressive::{ComponentCodecQuant, TILE_FLAG_DIFFERENCE};
use crate::dwt_extrapolate::BandInfo;
use crate::rlgr::RlgrError;
use crate::srl::{self, SrlError};
pub const COEFFICIENTS_PER_COMPONENT: usize = 4096;
type DecDwtQ = [[i16; COEFFICIENTS_PER_COMPONENT]; 3];
type SubBandDiffingTileKey = (u16, u16, u16);
pub const NUM_BANDS: usize = 10;
pub const SIGN_ZERO: i8 = 0;
pub const SIGN_POSITIVE: i8 = 1;
pub const SIGN_NEGATIVE: i8 = -1;
pub fn decode_first_pass(
data: &[u8],
base_quant: &ComponentCodecQuant,
prog_quant: &ComponentCodecQuant,
use_reduce_extrapolate: bool,
coefficients: &mut [i16],
sign: &mut [i8],
) -> Result<(), RlgrError> {
assert!(coefficients.len() >= COEFFICIENTS_PER_COMPONENT);
assert!(sign.len() >= COEFFICIENTS_PER_COMPONENT);
decode_first_pass_to_dwtq(data, prog_quant, use_reduce_extrapolate, coefficients, sign)?;
dequantize_component_ccq(coefficients, base_quant, use_reduce_extrapolate);
Ok(())
}
fn decode_first_pass_to_dwtq(
data: &[u8],
prog_quant: &ComponentCodecQuant,
use_reduce_extrapolate: bool,
coefficients: &mut [i16],
sign: &mut [i8],
) -> Result<(), RlgrError> {
crate::rlgr::decode(EntropyAlgorithm::Rlgr1, data, coefficients)?;
crate::subband_reconstruction::decode(&mut coefficients[ll3_offset(use_reduce_extrapolate)..]);
progressive_dequantize(coefficients, prog_quant, use_reduce_extrapolate);
capture_sign(coefficients, sign);
Ok(())
}
pub fn decode_upgrade_pass(
srl_data: &[u8],
raw_data: &[u8],
prev_prog_quant: &ComponentCodecQuant,
curr_prog_quant: &ComponentCodecQuant,
use_reduce_extrapolate: bool,
coefficients: &mut [i16],
sign: &mut [i8],
) -> Result<(), SrlError> {
assert!(coefficients.len() >= COEFFICIENTS_PER_COMPONENT);
assert!(sign.len() >= COEFFICIENTS_PER_COMPONENT);
let bands = get_band_layout(use_reduce_extrapolate);
let zero_counts: [usize; NUM_BANDS] = core::array::from_fn(|band_idx| band_zero_count(sign, &bands[band_idx]));
let has_srl_values = bands.iter().enumerate().any(|(band_idx, _)| {
let num_bits = prev_prog_quant
.for_band(band_idx)
.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_values = Vec::with_capacity(NUM_BANDS);
for (band_idx, _) in bands.iter().enumerate() {
let prev_bit_pos = prev_prog_quant.for_band(band_idx);
let curr_bit_pos = curr_prog_quant.for_band(band_idx);
let num_bits = prev_bit_pos.saturating_sub(curr_bit_pos);
if num_bits == 0 {
srl_values.push(Vec::new());
continue;
}
if band_idx == NUM_BANDS - 1 {
srl_values.push(Vec::new());
continue;
}
let zero_count = zero_counts[band_idx];
let values = match srl_decoder.as_mut() {
Some(decoder) => decoder.decode(zero_count, num_bits)?,
None => Vec::new(),
};
srl_values.push(values);
}
let mut raw_reader = RawBitReader::new(raw_data);
for (band_idx, band) in bands.iter().enumerate() {
let prev_bit_pos = prev_prog_quant.for_band(band_idx);
let curr_bit_pos = curr_prog_quant.for_band(band_idx);
let num_bits = prev_bit_pos.saturating_sub(curr_bit_pos);
if num_bits == 0 {
continue;
}
let is_ll3 = band_idx == NUM_BANDS - 1;
let mut srl_idx = 0;
for i in 0..band.count() {
let coeff_idx = band.offset + i;
if !is_ll3 && sign[coeff_idx] == SIGN_ZERO {
let value = srl_values[band_idx][srl_idx];
srl_idx += 1;
if value != 0 {
let shifted = i32::from(value) << i32::from(curr_bit_pos);
coefficients[coeff_idx] = clamp_i16(i32::from(coefficients[coeff_idx]) + shifted);
sign[coeff_idx] = if value > 0 { SIGN_POSITIVE } else { SIGN_NEGATIVE };
}
} else {
let raw_mag = raw_reader.read_bits(u32::from(num_bits));
if raw_mag != 0 {
let mag_i32 = i32::try_from(raw_mag).unwrap_or(i32::MAX);
let shifted = mag_i32 << i32::from(curr_bit_pos);
if is_ll3 || sign[coeff_idx] == SIGN_POSITIVE {
coefficients[coeff_idx] = clamp_i16(i32::from(coefficients[coeff_idx]) + shifted);
} else {
coefficients[coeff_idx] = clamp_i16(i32::from(coefficients[coeff_idx]) - shifted);
}
}
}
}
}
Ok(())
}
fn progressive_dequantize(coefficients: &mut [i16], prog_quant: &ComponentCodecQuant, use_reduce_extrapolate: bool) {
let bands = get_band_layout(use_reduce_extrapolate);
for (band_idx, band) in bands.iter().enumerate() {
let bit_pos = prog_quant.for_band(band_idx);
if bit_pos == 0 {
continue;
}
let is_ll3 = band_idx == 9;
let start = band.offset;
let end = start + band.count();
if is_ll3 {
for coeff in &mut coefficients[start..end] {
*coeff = clamp_i16(i32::from(*coeff) << i32::from(bit_pos));
}
} else {
for coeff in &mut coefficients[start..end] {
let val = i32::from(*coeff);
if val >= 0 {
*coeff = clamp_i16(val << i32::from(bit_pos));
} else {
*coeff = clamp_i16(-((-val) << i32::from(bit_pos)));
}
}
}
}
}
pub fn progressive_quantize(coefficients: &mut [i16], prog_quant: &ComponentCodecQuant, use_reduce_extrapolate: bool) {
let bands = get_band_layout(use_reduce_extrapolate);
for (band_idx, band) in bands.iter().enumerate() {
let bit_pos = prog_quant.for_band(band_idx);
if bit_pos == 0 {
continue;
}
let is_ll3 = band_idx == 9;
let start = band.offset;
let end = start + band.count();
if is_ll3 {
for coeff in &mut coefficients[start..end] {
*coeff >>= bit_pos;
}
} else {
for coeff in &mut coefficients[start..end] {
let val = i32::from(*coeff);
if val >= 0 {
*coeff = clamp_i16(val >> i32::from(bit_pos));
} else {
*coeff = clamp_i16(-((-val) >> i32::from(bit_pos)));
}
}
}
}
}
pub fn encode_first_pass(
coefficients: &mut [i16],
output: &mut [u8],
base_quant: &ComponentCodecQuant,
prog_quant: &ComponentCodecQuant,
use_reduce_extrapolate: bool,
) -> Result<usize, RlgrError> {
assert!(coefficients.len() >= COEFFICIENTS_PER_COMPONENT);
let mut temp = [0i16; COEFFICIENTS_PER_COMPONENT];
if use_reduce_extrapolate {
crate::dwt_extrapolate::encode(coefficients, &mut temp);
} else {
crate::dwt::encode(coefficients, &mut temp);
}
quantize_component_ccq(coefficients, base_quant, use_reduce_extrapolate);
progressive_quantize(coefficients, prog_quant, use_reduce_extrapolate);
crate::subband_reconstruction::encode(&mut coefficients[ll3_offset(use_reduce_extrapolate)..]);
crate::rlgr::encode(EntropyAlgorithm::Rlgr1, coefficients, output)
}
fn quantize_component_ccq(coefficients: &mut [i16], quant: &ComponentCodecQuant, use_reduce_extrapolate: bool) {
let bands = get_band_layout(use_reduce_extrapolate);
for (band_idx, band) in bands.iter().enumerate() {
let q = quant.for_band(band_idx);
let start = band.offset;
let end = start + band.count();
match q.cmp(&6) {
core::cmp::Ordering::Greater => {
let shift = q - 6;
for coeff in &mut coefficients[start..end] {
*coeff = round_shift_right(*coeff, shift);
}
}
core::cmp::Ordering::Less => {
let shift = 6 - q;
for coeff in &mut coefficients[start..end] {
*coeff = clamp_i16(i32::from(*coeff) << i32::from(shift));
}
}
core::cmp::Ordering::Equal => {}
}
}
}
pub fn encode_upgrade_pass(
coefficients: &[i16],
prev_coefficients: &[i16],
prev_prog_quant: &ComponentCodecQuant,
curr_prog_quant: &ComponentCodecQuant,
sign: &[i8],
use_reduce_extrapolate: bool,
) -> Result<(Vec<u8>, Vec<u8>), SrlError> {
let bands = get_band_layout(use_reduce_extrapolate);
let mut srl_encoder = srl::SrlEncoder::new();
let mut has_srl_values = false;
let mut raw_writer = RawBitWriter::new();
for (band_idx, band) in bands.iter().enumerate() {
let prev_bit_pos = prev_prog_quant.for_band(band_idx);
let curr_bit_pos = curr_prog_quant.for_band(band_idx);
let num_bits = prev_bit_pos.saturating_sub(curr_bit_pos);
if num_bits == 0 {
continue;
}
let mut band_srl_values = Vec::new();
for i in 0..band.count() {
let coeff_idx = band.offset + i;
let is_ll3 = band_idx == NUM_BANDS - 1;
if !is_ll3 && sign[coeff_idx] == SIGN_ZERO {
let curr_shifted = i32::from(coefficients[coeff_idx]) >> i32::from(curr_bit_pos);
let prev_shifted = i32::from(prev_coefficients[coeff_idx]) >> i32::from(curr_bit_pos);
let delta = clamp_i16(curr_shifted - prev_shifted);
band_srl_values.push(delta);
} else {
let curr_abs = i32::from(coefficients[coeff_idx]).unsigned_abs();
let prev_abs = i32::from(prev_coefficients[coeff_idx]).unsigned_abs();
let curr_q = curr_abs >> u32::from(curr_bit_pos);
let prev_q = prev_abs >> u32::from(curr_bit_pos);
let raw_mag = curr_q.saturating_sub(prev_q);
raw_writer.write_bits(raw_mag, u32::from(num_bits));
}
}
if !band_srl_values.is_empty() {
srl_encoder.encode(&band_srl_values, num_bits)?;
has_srl_values = true;
}
}
let raw_data = raw_writer.finish();
let srl_data = if has_srl_values {
srl_encoder.finish()?
} else {
Vec::new()
};
Ok((srl_data, raw_data))
}
#[expect(clippy::similar_names)]
pub fn rgba_to_ycbcr(pixels: &[u8], y_out: &mut [i16], cb_out: &mut [i16], cr_out: &mut [i16]) {
assert!(pixels.len() >= 64 * 64 * 4);
assert!(y_out.len() >= COEFFICIENTS_PER_COMPONENT);
assert!(cb_out.len() >= COEFFICIENTS_PER_COMPONENT);
assert!(cr_out.len() >= COEFFICIENTS_PER_COMPONENT);
for i in 0..64 * 64 {
let off = i * 4;
let r = i32::from(pixels[off]);
let g = i32::from(pixels[off + 1]);
let b = i32::from(pixels[off + 2]);
let y = ((19595 * r + 38470 * g + 7471 * b + 32768) >> 16) - 128;
let cb = (-11059 * r - 21709 * g + 32768 * b + 32768) >> 16;
let cr = (32768 * r - 27439 * g - 5329 * b + 32768) >> 16;
y_out[i] = clamp_i16(y);
cb_out[i] = clamp_i16(cb);
cr_out[i] = clamp_i16(cr);
}
}
fn dequantize_component_ccq(coefficients: &mut [i16], quant: &ComponentCodecQuant, use_reduce_extrapolate: bool) {
let bands = get_band_layout(use_reduce_extrapolate);
for (band_idx, band) in bands.iter().enumerate() {
let q = quant.for_band(band_idx);
let start = band.offset;
let end = start + band.count();
match q.cmp(&6) {
core::cmp::Ordering::Greater => {
let shift = q - 6;
for coeff in &mut coefficients[start..end] {
*coeff = clamp_i16(i32::from(*coeff) << i32::from(shift));
}
}
core::cmp::Ordering::Less => {
let shift = 6 - q;
for coeff in &mut coefficients[start..end] {
*coeff = round_shift_right(*coeff, shift);
}
}
core::cmp::Ordering::Equal => {}
}
}
}
fn capture_sign(coefficients: &[i16], sign: &mut [i8]) {
for (s, &c) in sign.iter_mut().zip(coefficients.iter()) {
*s = match c.cmp(&0) {
core::cmp::Ordering::Greater => SIGN_POSITIVE,
core::cmp::Ordering::Less => SIGN_NEGATIVE,
core::cmp::Ordering::Equal => SIGN_ZERO,
};
}
}
fn get_band_layout(use_reduce_extrapolate: bool) -> [BandInfo; NUM_BANDS] {
if use_reduce_extrapolate {
crate::dwt_extrapolate::band_layout()
} else {
standard_band_layout()
}
}
fn standard_band_layout() -> [BandInfo; NUM_BANDS] {
let mut off = 0;
let mut b = |w: usize, h: usize| {
let info = BandInfo {
width: w,
height: h,
offset: off,
};
off += w * h;
info
};
[
b(32, 32), b(32, 32), b(32, 32), b(16, 16), b(16, 16), b(16, 16), b(8, 8), b(8, 8), b(8, 8), b(8, 8), ]
}
fn ll3_offset(use_reduce_extrapolate: bool) -> usize {
if use_reduce_extrapolate {
4015 } else {
4032 }
}
fn band_zero_count(sign: &[i8], band: &BandInfo) -> usize {
let start = band.offset;
let end = start + band.count();
sign[start..end].iter().filter(|&&s| s == SIGN_ZERO).count()
}
#[expect(
clippy::as_conversions,
clippy::cast_sign_loss,
reason = "value is clamped to 0..255 before cast"
)]
fn clamp_u8(value: i32) -> u8 {
value.clamp(0, 255) as u8
}
#[expect(
clippy::as_conversions,
clippy::cast_possible_truncation,
reason = "value is clamped to i16 range before cast"
)]
fn clamp_i16(value: i32) -> i16 {
value.clamp(i32::from(i16::MIN), i32::from(i16::MAX)) as i16
}
fn round_shift_right(value: i16, shift: u8) -> i16 {
debug_assert!(shift > 0);
let half = 1i32 << (i32::from(shift) - 1);
clamp_i16((i32::from(value) + half) >> i32::from(shift))
}
struct RawBitWriter {
bytes: Vec<u8>,
current: u8,
bit_count: u8,
}
impl RawBitWriter {
fn new() -> Self {
Self {
bytes: Vec::new(),
current: 0,
bit_count: 0,
}
}
fn write_bit(&mut self, bit: bool) {
self.current = (self.current << 1) | u8::from(bit);
self.bit_count += 1;
if self.bit_count >= 8 {
self.bytes.push(self.current);
self.current = 0;
self.bit_count = 0;
}
}
fn write_bits(&mut self, value: u32, count: u32) {
debug_assert!(count <= 32, "RawBitWriter::write_bits count must be <= 32");
for i in (0..count).rev() {
self.write_bit((value >> i) & 1 != 0);
}
}
fn finish(mut self) -> Vec<u8> {
if self.bit_count > 0 {
self.current <<= 8 - self.bit_count;
self.bytes.push(self.current);
}
self.bytes
}
}
struct RawBitReader<'a> {
data: &'a [u8],
byte_idx: usize,
bit_idx: u8,
}
impl<'a> RawBitReader<'a> {
fn new(data: &'a [u8]) -> Self {
Self {
data,
byte_idx: 0,
bit_idx: 0,
}
}
fn read_bits(&mut self, count: u32) -> u32 {
let mut value = 0u32;
for _ in 0..count {
value = (value << 1) | u32::from(self.read_bit());
}
value
}
fn read_bit(&mut self) -> bool {
if self.byte_idx >= self.data.len() {
return false;
}
let bit = (self.data[self.byte_idx] >> (7 - self.bit_idx)) & 1 != 0;
self.bit_idx += 1;
if self.bit_idx >= 8 {
self.bit_idx = 0;
self.byte_idx += 1;
}
bit
}
}
pub struct TileState {
pub coefficients: DecDwtQ,
pub sign: [[i8; COEFFICIENTS_PER_COMPONENT]; 3],
pub prog_quant: [ComponentCodecQuant; 3],
pub quant_idx: [u8; 3],
pub base_quant: [ComponentCodecQuant; 3],
pub pass: u16,
pub is_difference: bool,
pub quality: u8,
pub use_reduce_extrapolate: bool,
}
struct FirstPassOptions {
quant_idx: [u8; 3],
quality: u8,
use_reduce_extrapolate: bool,
}
impl TileState {
pub fn new() -> Self {
Self {
coefficients: [[0; COEFFICIENTS_PER_COMPONENT]; 3],
sign: [[0; COEFFICIENTS_PER_COMPONENT]; 3],
prog_quant: [ComponentCodecQuant::LOSSLESS; 3],
quant_idx: [0; 3],
base_quant: [ComponentCodecQuant {
ll3: 6,
hl3: 6,
lh3: 6,
hh3: 6,
hl2: 6,
lh2: 6,
hh2: 6,
hl1: 6,
lh1: 6,
hh1: 6,
}; 3],
pass: 0,
is_difference: false,
quality: 0,
use_reduce_extrapolate: false,
}
}
pub fn decode_first(
&mut self,
component_data: [&[u8]; 3],
base_quants: [&ComponentCodecQuant; 3],
prog_quants: [ComponentCodecQuant; 3],
quant_idx: [u8; 3],
quality: u8,
use_reduce_extrapolate: bool,
) -> Result<(), RlgrError> {
self.decode_first_with_difference(
component_data,
base_quants,
prog_quants,
None,
FirstPassOptions {
quant_idx,
quality,
use_reduce_extrapolate,
},
)
}
fn decode_first_with_difference(
&mut self,
component_data: [&[u8]; 3],
base_quants: [&ComponentCodecQuant; 3],
prog_quants: [ComponentCodecQuant; 3],
reference: Option<&DecDwtQ>,
options: FirstPassOptions,
) -> Result<(), RlgrError> {
let mut coefficients = [[0; COEFFICIENTS_PER_COMPONENT]; 3];
let mut sign = [[SIGN_ZERO; COEFFICIENTS_PER_COMPONENT]; 3];
for c in 0..3 {
decode_first_pass_to_dwtq(
component_data[c],
&prog_quants[c],
options.use_reduce_extrapolate,
&mut coefficients[c],
&mut sign[c],
)?;
}
if let Some(reference) = reference {
for (component, reference_component) in coefficients.iter_mut().zip(reference.iter()) {
for (coefficient, reference_coefficient) in component.iter_mut().zip(reference_component.iter()) {
*coefficient = coefficient.saturating_add(*reference_coefficient);
}
}
}
self.coefficients = coefficients;
self.sign = sign;
self.pass = 1;
self.quality = options.quality;
self.quant_idx = options.quant_idx;
self.base_quant = [*base_quants[0], *base_quants[1], *base_quants[2]];
self.use_reduce_extrapolate = options.use_reduce_extrapolate;
self.is_difference = reference.is_some();
self.prog_quant = prog_quants;
Ok(())
}
pub fn decode_upgrade(
&mut self,
srl_data: [&[u8]; 3],
raw_data: [&[u8]; 3],
prog_quants: [ComponentCodecQuant; 3],
quality: u8,
) -> Result<(), SrlError> {
let prev_prog_quant = self.prog_quant;
let mut coefficients = self.coefficients;
let mut sign = self.sign;
for c in 0..3 {
decode_upgrade_pass(
srl_data[c],
raw_data[c],
&prev_prog_quant[c],
&prog_quants[c],
self.use_reduce_extrapolate,
&mut coefficients[c],
&mut sign[c],
)?;
}
self.coefficients = coefficients;
self.sign = sign;
self.prog_quant = prog_quants;
self.quality = quality;
self.pass = self.pass.saturating_add(1);
Ok(())
}
#[expect(clippy::similar_names, reason = "y/cb/cr are standard YCbCr component names")]
pub fn reconstruct_to_rgba(&self, pixels: &mut [u8]) {
assert!(pixels.len() >= 64 * 64 * 4, "pixel buffer too small");
let mut y_buf = self.coefficients[0];
let mut cb_buf = self.coefficients[1];
let mut cr_buf = self.coefficients[2];
let mut temp = [0i16; COEFFICIENTS_PER_COMPONENT];
dequantize_component_ccq(&mut y_buf, &self.base_quant[0], self.use_reduce_extrapolate);
dequantize_component_ccq(&mut cb_buf, &self.base_quant[1], self.use_reduce_extrapolate);
dequantize_component_ccq(&mut cr_buf, &self.base_quant[2], self.use_reduce_extrapolate);
if self.use_reduce_extrapolate {
crate::dwt_extrapolate::decode(&mut y_buf, &mut temp);
crate::dwt_extrapolate::decode(&mut cb_buf, &mut temp);
crate::dwt_extrapolate::decode(&mut cr_buf, &mut temp);
} else {
let mut dwt_temp = [0i16; COEFFICIENTS_PER_COMPONENT];
crate::dwt::decode(&mut y_buf, &mut dwt_temp);
crate::dwt::decode(&mut cb_buf, &mut dwt_temp);
crate::dwt::decode(&mut cr_buf, &mut dwt_temp);
}
for i in 0..64 * 64 {
let y = i32::from(y_buf[i]) + 128;
let cb = i32::from(cb_buf[i]);
let cr = i32::from(cr_buf[i]);
let r = y + ((cr * 91881 + 32768) >> 16);
let g = y - ((cb * 22554 + cr * 46802 + 32768) >> 16);
let b = y + ((cb * 116130 + 32768) >> 16);
let off = i * 4;
pixels[off] = clamp_u8(r);
pixels[off + 1] = clamp_u8(g);
pixels[off + 2] = clamp_u8(b);
pixels[off + 3] = 0xFF;
}
}
}
impl Default for TileState {
fn default() -> Self {
Self::new()
}
}
pub struct SurfaceTiles {
pub tiles_wide: u16,
pub tiles_high: u16,
pub use_reduce_extrapolate: bool,
pub tiles: Vec<Option<Box<TileState>>>,
}
impl SurfaceTiles {
pub fn new(
width_pixels: u16,
height_pixels: u16,
use_reduce_extrapolate: bool,
) -> Result<Self, ProgressiveDecodeError> {
if width_pixels > MAX_SURFACE_DIM || height_pixels > MAX_SURFACE_DIM {
return Err(ProgressiveDecodeError::SurfaceTooLarge {
width: width_pixels,
height: height_pixels,
});
}
let tiles_wide = width_pixels.div_ceil(64);
let tiles_high = height_pixels.div_ceil(64);
let count = usize::from(tiles_wide) * usize::from(tiles_high);
Ok(Self {
tiles_wide,
tiles_high,
use_reduce_extrapolate,
tiles: core::iter::repeat_with(|| None).take(count).collect(),
})
}
pub fn get_or_create(&mut self, x_idx: u16, y_idx: u16) -> Option<&mut TileState> {
let idx = self.tile_index(x_idx, y_idx)?;
let tile = self.tiles[idx].get_or_insert_with(|| {
let mut t = Box::new(TileState::new());
t.use_reduce_extrapolate = self.use_reduce_extrapolate;
t
});
Some(tile)
}
pub fn get(&self, x_idx: u16, y_idx: u16) -> Option<&TileState> {
let idx = self.tile_index(x_idx, y_idx)?;
self.tiles[idx].as_deref()
}
pub fn reset(&mut self) {
for tile in &mut self.tiles {
*tile = None;
}
}
fn tile_index(&self, x_idx: u16, y_idx: u16) -> Option<usize> {
if x_idx >= self.tiles_wide || y_idx >= self.tiles_high {
return None;
}
Some(usize::from(y_idx) * usize::from(self.tiles_wide) + usize::from(x_idx))
}
}
pub struct DecodedTile {
pub x_idx: u16,
pub y_idx: u16,
pub pixels: Vec<u8>,
}
pub const MAX_SURFACE_DIM: u16 = 32768;
#[derive(Debug)]
pub enum ProgressiveDecodeError {
Pdu(rat_rdp_core::DecodeError),
Rlgr(RlgrError),
Srl(SrlError),
MissingBlock(&'static str),
TileOutOfBounds { x_idx: u16, y_idx: u16 },
InvalidQuantIndex { index: usize, table_len: usize },
MissingTileReference { x_idx: u16, y_idx: u16 },
SurfaceTooLarge { width: u16, height: u16 },
}
impl core::fmt::Display for ProgressiveDecodeError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Pdu(e) => write!(f, "progressive PDU decode: {e}"),
Self::Rlgr(e) => write!(f, "progressive RLGR decode: {e}"),
Self::Srl(e) => write!(f, "progressive srl decode: {e}"),
Self::MissingBlock(name) => write!(f, "progressive stream missing {name} block"),
Self::TileOutOfBounds { x_idx, y_idx } => {
write!(f, "tile ({x_idx}, {y_idx}) out of surface bounds")
}
Self::InvalidQuantIndex { index, table_len } => {
write!(f, "quant index {index} exceeds table length {table_len}")
}
Self::MissingTileReference { x_idx, y_idx } => {
write!(f, "difference tile ({x_idx}, {y_idx}) has no retained reference")
}
Self::SurfaceTooLarge { width, height } => {
write!(
f,
"surface dimensions {width}x{height} exceed per-axis cap of {MAX_SURFACE_DIM} \
(MS-RDPEGFX 2.2.2.14 normative ceiling: 32766)"
)
}
}
}
}
impl From<rat_rdp_core::DecodeError> for ProgressiveDecodeError {
fn from(e: rat_rdp_core::DecodeError) -> Self {
Self::Pdu(e)
}
}
impl From<RlgrError> for ProgressiveDecodeError {
fn from(e: RlgrError) -> Self {
Self::Rlgr(e)
}
}
impl From<SrlError> for ProgressiveDecodeError {
fn from(e: SrlError) -> Self {
Self::Srl(e)
}
}
struct ProgressiveContext {
surface: SurfaceTiles,
}
pub struct ProgressiveDecoder {
contexts: BTreeMap<(u16, u32), ProgressiveContext>,
references: BTreeMap<SubBandDiffingTileKey, DecDwtQ>,
}
impl ProgressiveDecoder {
pub fn new() -> Self {
Self {
contexts: BTreeMap::new(),
references: BTreeMap::new(),
}
}
pub fn decode_bitmap(
&mut self,
surface_id: u16,
codec_context_id: u32,
surface_width: u16,
surface_height: u16,
bitmap_data: &[u8],
) -> Result<Vec<DecodedTile>, ProgressiveDecodeError> {
use rat_rdp_pdu::codecs::rfx::progressive::{ProgressiveBlock, decode_progressive_stream};
let blocks = decode_progressive_stream(bitmap_data)?;
let use_reduce_extrapolate = match blocks.iter().find_map(|block| match block {
ProgressiveBlock::Context(ctx) => Some(ctx.uses_reduce_extrapolate()),
_ => None,
}) {
Some(v) => v,
None => self
.contexts
.get(&(surface_id, codec_context_id))
.map(|c| c.surface.use_reduce_extrapolate)
.ok_or(ProgressiveDecodeError::MissingBlock("CONTEXT"))?,
};
let (contexts, references) = (&mut self.contexts, &mut self.references);
let context = match contexts.entry((surface_id, codec_context_id)) {
Entry::Occupied(e) => e.into_mut(),
Entry::Vacant(e) => {
let surface = SurfaceTiles::new(surface_width, surface_height, use_reduce_extrapolate)?;
e.insert(ProgressiveContext { surface })
}
};
let expected_wide = surface_width.div_ceil(64);
let expected_high = surface_height.div_ceil(64);
if context.surface.tiles_wide != expected_wide || context.surface.tiles_high != expected_high {
context.surface = SurfaceTiles::new(surface_width, surface_height, use_reduce_extrapolate)?;
}
context.surface.use_reduce_extrapolate = use_reduce_extrapolate;
let mut decoded_tiles = Vec::new();
for block in &blocks {
let region = match block {
ProgressiveBlock::Region(r) => r,
_ => continue,
};
let quant_vals = ®ion.quant_vals;
let prog_quant_vals = ®ion.quant_prog_vals;
for tile_block in ®ion.tiles {
let tiles = decode_tile_block(
surface_id,
&mut context.surface,
references,
tile_block,
quant_vals,
prog_quant_vals,
use_reduce_extrapolate,
)?;
decoded_tiles.extend(tiles);
}
}
Ok(decoded_tiles)
}
pub fn delete_context(&mut self, surface_id: u16, codec_context_id: u32) {
self.contexts.remove(&(surface_id, codec_context_id));
}
pub fn delete_surface(&mut self, surface_id: u16) {
self.contexts
.retain(|(context_surface_id, _), _| *context_surface_id != surface_id);
self.references
.retain(|(reference_surface_id, _, _), _| *reference_surface_id != surface_id);
}
pub fn reset(&mut self) {
self.contexts.clear();
}
}
#[expect(
clippy::similar_names,
reason = "q_y/q_cb/q_cr are standard component quant index names"
)]
fn decode_tile_block(
surface_id: u16,
surface: &mut SurfaceTiles,
references: &mut BTreeMap<SubBandDiffingTileKey, DecDwtQ>,
tile_block: &rat_rdp_pdu::codecs::rfx::progressive::ProgressiveTile<'_>,
quant_vals: &[ComponentCodecQuant],
prog_quant_vals: &[rat_rdp_pdu::codecs::rfx::progressive::ProgressiveCodecQuant],
use_reduce_extrapolate: bool,
) -> Result<Vec<DecodedTile>, ProgressiveDecodeError> {
use rat_rdp_pdu::codecs::rfx::progressive::ProgressiveTile;
match tile_block {
ProgressiveTile::Simple(tile) => {
let x_idx = tile.x_idx;
let y_idx = tile.y_idx;
let is_difference = tile.flags & TILE_FLAG_DIFFERENCE != 0;
if surface.tile_index(x_idx, y_idx).is_none() {
return Err(ProgressiveDecodeError::TileOutOfBounds { x_idx, y_idx });
}
let reference_key = (surface_id, x_idx, y_idx);
let reference = if is_difference {
Some(
references
.get(&reference_key)
.ok_or(ProgressiveDecodeError::MissingTileReference { x_idx, y_idx })?,
)
} else {
None
};
let tile_state = surface
.get_or_create(x_idx, y_idx)
.ok_or(ProgressiveDecodeError::TileOutOfBounds { x_idx, y_idx })?;
let q_y = usize::from(tile.quant_idx_y);
let q_cb = usize::from(tile.quant_idx_cb);
let q_cr = usize::from(tile.quant_idx_cr);
if q_y >= quant_vals.len() || q_cb >= quant_vals.len() || q_cr >= quant_vals.len() {
return Err(ProgressiveDecodeError::InvalidQuantIndex {
index: q_y.max(q_cb).max(q_cr),
table_len: quant_vals.len(),
});
}
let prog = ComponentCodecQuant::LOSSLESS;
tile_state.decode_first_with_difference(
[tile.y_data, tile.cb_data, tile.cr_data],
[&quant_vals[q_y], &quant_vals[q_cb], &quant_vals[q_cr]],
[prog, prog, prog],
reference,
FirstPassOptions {
quant_idx: [tile.quant_idx_y, tile.quant_idx_cb, tile.quant_idx_cr],
quality: 0xFF, use_reduce_extrapolate,
},
)?;
references.insert(reference_key, tile_state.coefficients);
let mut pixels = vec![0u8; 64 * 64 * 4];
tile_state.reconstruct_to_rgba(&mut pixels);
Ok(vec![DecodedTile { x_idx, y_idx, pixels }])
}
ProgressiveTile::First(tile) => {
let x_idx = tile.x_idx;
let y_idx = tile.y_idx;
let is_difference = tile.flags & TILE_FLAG_DIFFERENCE != 0;
if surface.tile_index(x_idx, y_idx).is_none() {
return Err(ProgressiveDecodeError::TileOutOfBounds { x_idx, y_idx });
}
let reference_key = (surface_id, x_idx, y_idx);
let reference = if is_difference {
Some(
references
.get(&reference_key)
.ok_or(ProgressiveDecodeError::MissingTileReference { x_idx, y_idx })?,
)
} else {
None
};
let tile_state = surface
.get_or_create(x_idx, y_idx)
.ok_or(ProgressiveDecodeError::TileOutOfBounds { x_idx, y_idx })?;
let q_y = usize::from(tile.quant_idx_y);
let q_cb = usize::from(tile.quant_idx_cb);
let q_cr = usize::from(tile.quant_idx_cr);
if q_y >= quant_vals.len() || q_cb >= quant_vals.len() || q_cr >= quant_vals.len() {
return Err(ProgressiveDecodeError::InvalidQuantIndex {
index: q_y.max(q_cb).max(q_cr),
table_len: quant_vals.len(),
});
}
let pq_idx = usize::from(tile.quality);
if pq_idx >= prog_quant_vals.len() {
return Err(ProgressiveDecodeError::InvalidQuantIndex {
index: pq_idx,
table_len: prog_quant_vals.len(),
});
}
let pq = &prog_quant_vals[pq_idx];
tile_state.decode_first_with_difference(
[tile.y_data, tile.cb_data, tile.cr_data],
[&quant_vals[q_y], &quant_vals[q_cb], &quant_vals[q_cr]],
[pq.y_quant, pq.cb_quant, pq.cr_quant],
reference,
FirstPassOptions {
quant_idx: [tile.quant_idx_y, tile.quant_idx_cb, tile.quant_idx_cr],
quality: tile.quality,
use_reduce_extrapolate,
},
)?;
references.insert(reference_key, tile_state.coefficients);
let mut pixels = vec![0u8; 64 * 64 * 4];
tile_state.reconstruct_to_rgba(&mut pixels);
Ok(vec![DecodedTile { x_idx, y_idx, pixels }])
}
ProgressiveTile::Upgrade(tile) => {
let x_idx = tile.x_idx;
let y_idx = tile.y_idx;
let tile_state = surface
.get_or_create(x_idx, y_idx)
.ok_or(ProgressiveDecodeError::TileOutOfBounds { x_idx, y_idx })?;
if tile_state.pass == 0 {
return Ok(Vec::new());
}
let pq_idx = usize::from(tile.quality);
if pq_idx >= prog_quant_vals.len() {
return Err(ProgressiveDecodeError::InvalidQuantIndex {
index: pq_idx,
table_len: prog_quant_vals.len(),
});
}
let pq = &prog_quant_vals[pq_idx];
tile_state.decode_upgrade(
[tile.y_srl_data, tile.cb_srl_data, tile.cr_srl_data],
[tile.y_raw_data, tile.cb_raw_data, tile.cr_raw_data],
[pq.y_quant, pq.cb_quant, pq.cr_quant],
tile.quality,
)?;
references.insert((surface_id, x_idx, y_idx), tile_state.coefficients);
let mut pixels = vec![0u8; 64 * 64 * 4];
tile_state.reconstruct_to_rgba(&mut pixels);
Ok(vec![DecodedTile { x_idx, y_idx, pixels }])
}
}
}
impl Default for ProgressiveDecoder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
#[expect(clippy::as_conversions, clippy::cast_possible_truncation, clippy::cast_possible_wrap)]
mod tests {
use super::*;
fn minimal_progressive_stream(include_context: bool) -> Vec<u8> {
use rat_rdp_pdu::codecs::rfx::RfxRectangle;
use rat_rdp_pdu::codecs::rfx::progressive::{
ProgressiveBlock, ProgressiveContextPdu, ProgressiveFrameBeginPdu, ProgressiveFrameEndPdu,
ProgressiveRegion, ProgressiveSyncPdu, encode_progressive_stream,
};
let region = ProgressiveRegion {
tile_size: 0x40,
rects: vec![RfxRectangle {
x: 0,
y: 0,
width: 64,
height: 64,
}],
quant_vals: vec![],
quant_prog_vals: vec![],
flags: 0,
tiles: vec![],
};
let mut blocks = vec![ProgressiveBlock::Sync(ProgressiveSyncPdu)];
if include_context {
blocks.push(ProgressiveBlock::Context(ProgressiveContextPdu {
context_id: 0,
tile_size: 0x0040,
flags: 0,
}));
}
blocks.extend([
ProgressiveBlock::FrameBegin(ProgressiveFrameBeginPdu {
frame_index: 0,
region_count: 1,
}),
ProgressiveBlock::Region(region),
ProgressiveBlock::FrameEnd(ProgressiveFrameEndPdu),
]);
encode_progressive_stream(&blocks).unwrap()
}
#[test]
fn surface_tiles_rejects_over_cap_dimensions() {
assert!(SurfaceTiles::new(MAX_SURFACE_DIM, MAX_SURFACE_DIM, false).is_ok());
let over_w = MAX_SURFACE_DIM.checked_add(1).unwrap();
match SurfaceTiles::new(over_w, 1024, false) {
Err(ProgressiveDecodeError::SurfaceTooLarge { width, height }) => {
assert_eq!(width, over_w);
assert_eq!(height, 1024);
}
Err(other) => panic!("expected SurfaceTooLarge, got Err({other})"),
Ok(_) => panic!("expected SurfaceTooLarge, got Ok"),
}
match SurfaceTiles::new(1024, over_w, false) {
Err(ProgressiveDecodeError::SurfaceTooLarge { width, height }) => {
assert_eq!(width, 1024);
assert_eq!(height, over_w);
}
Err(other) => panic!("expected SurfaceTooLarge, got Err({other})"),
Ok(_) => panic!("expected SurfaceTooLarge, got Ok"),
}
}
#[test]
fn standard_band_layout_totals_4096() {
let bands = standard_band_layout();
let total: usize = bands.iter().map(|b| b.count()).sum();
assert_eq!(total, 4096);
}
#[test]
fn standard_band_offsets() {
let bands = standard_band_layout();
assert_eq!(bands[0].offset, 0);
assert_eq!(bands[1].offset, 1024);
assert_eq!(bands[2].offset, 2048);
assert_eq!(bands[3].offset, 3072);
assert_eq!(bands[4].offset, 3328);
assert_eq!(bands[5].offset, 3584);
assert_eq!(bands[6].offset, 3840);
assert_eq!(bands[7].offset, 3904);
assert_eq!(bands[8].offset, 3968);
assert_eq!(bands[9].offset, 4032);
}
#[test]
fn sign_capture_tri_state() {
let coefficients = [10i16, -5, 0, 100, -1, 0];
let mut sign = [0i8; 6];
capture_sign(&coefficients, &mut sign);
assert_eq!(sign, [1, -1, 0, 1, -1, 0]);
}
#[test]
fn progressive_dequantize_ll3_shift() {
let mut coefficients = vec![0i16; 4096];
coefficients[4032] = 5;
coefficients[4033] = -3;
let prog_quant = ComponentCodecQuant {
ll3: 2,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 0,
lh1: 0,
hh1: 0,
};
progressive_dequantize(&mut coefficients, &prog_quant, false);
assert_eq!(coefficients[4032], 20);
assert_eq!(coefficients[4033], -12);
}
#[test]
fn progressive_dequantize_non_ll3_preserves_sign() {
let mut coefficients = vec![0i16; 4096];
coefficients[0] = 5;
coefficients[1] = -5;
let prog_quant = ComponentCodecQuant {
ll3: 0,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 2,
lh1: 0,
hh1: 0,
};
progressive_dequantize(&mut coefficients, &prog_quant, false);
assert_eq!(coefficients[0], 20); assert_eq!(coefficients[1], -20); }
#[test]
fn progressive_quantize_round_trip() {
let mut coefficients = vec![0i16; 4096];
for (i, c) in coefficients.iter_mut().enumerate() {
*c = (i as i16).wrapping_mul(7);
}
let original = coefficients.clone();
let prog_quant = ComponentCodecQuant {
ll3: 2,
hl3: 3,
lh3: 3,
hh3: 4,
hl2: 3,
lh2: 3,
hh2: 4,
hl1: 2,
lh1: 2,
hh1: 3,
};
progressive_quantize(&mut coefficients, &prog_quant, false);
progressive_dequantize(&mut coefficients, &prog_quant, false);
for (i, (&a, &b)) in coefficients.iter().zip(original.iter()).enumerate() {
let err = (i32::from(a) - i32::from(b)).unsigned_abs();
assert!(err < 32, "index {i}: error {err} too large");
}
}
#[test]
fn raw_bit_reader_basic() {
let data = [0b10110000, 0b01010000];
let mut reader = RawBitReader::new(&data);
assert_eq!(reader.read_bits(4), 0b1011);
assert_eq!(reader.read_bits(4), 0b0000);
assert_eq!(reader.read_bits(4), 0b0101);
}
#[test]
fn clamp_i16_limits() {
assert_eq!(clamp_i16(40000), i16::MAX);
assert_eq!(clamp_i16(-40000), i16::MIN);
assert_eq!(clamp_i16(100), 100);
assert_eq!(clamp_i16(-100), -100);
}
#[test]
fn band_zero_count_counts_correctly() {
let mut sign = [0i8; 4096];
sign[0] = SIGN_POSITIVE;
sign[1] = SIGN_NEGATIVE;
sign[2] = SIGN_ZERO;
let bands = standard_band_layout();
assert_eq!(band_zero_count(&sign, &bands[0]), 1022); }
#[test]
fn ll3_offsets_correct() {
assert_eq!(ll3_offset(false), 4032);
assert_eq!(ll3_offset(true), 4015);
}
#[test]
fn upgrade_pass_zero_das_becomes_nonzero() {
let mut coefficients = vec![0i16; 4096];
let mut sign = vec![SIGN_POSITIVE; 4096];
sign[0] = SIGN_ZERO;
let prev_prog_quant = ComponentCodecQuant {
ll3: 0,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 4,
lh1: 0,
hh1: 0,
};
let curr_prog_quant = ComponentCodecQuant {
ll3: 0,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 2,
lh1: 0,
hh1: 0,
};
let srl_data = vec![0b1001_0000, 0x00];
let raw_data = vec![];
decode_upgrade_pass(
&srl_data,
&raw_data,
&prev_prog_quant,
&curr_prog_quant,
false,
&mut coefficients,
&mut sign,
)
.unwrap();
assert_eq!(coefficients[0], 4);
assert_eq!(sign[0], SIGN_POSITIVE);
}
#[test]
fn upgrade_pass_preserves_component_streams_across_bands() {
let mut coefficients = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut sign = [SIGN_POSITIVE; COEFFICIENTS_PER_COMPONENT];
sign[0] = SIGN_ZERO;
sign[1024] = SIGN_ZERO;
let mut prev_prog_quant = ComponentCodecQuant::LOSSLESS;
prev_prog_quant.hl1 = 1;
prev_prog_quant.lh1 = 1;
decode_upgrade_pass(
&[0b1001_1000, 0x00],
&[0b1000_0000],
&prev_prog_quant,
&ComponentCodecQuant::LOSSLESS,
false,
&mut coefficients,
&mut sign,
)
.unwrap();
assert_eq!(coefficients[0], 1);
assert_eq!(coefficients[1024], -1);
assert_eq!(coefficients[1], 1);
assert_eq!(coefficients[1025], 0);
}
#[test]
fn upgrade_pass_rejects_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;
assert_eq!(
decode_upgrade_pass(
&[0x80, 0x00],
&[],
&prev_prog_quant,
&ComponentCodecQuant::LOSSLESS,
false,
&mut coefficients,
&mut sign,
),
Err(SrlError::Truncated)
);
}
#[test]
fn tile_upgrade_keeps_all_components_on_srl_error() {
let mut tile = TileState::new();
let mut prev_prog_quant = ComponentCodecQuant::LOSSLESS;
prev_prog_quant.hl1 = 4;
tile.prog_quant = [prev_prog_quant; 3];
tile.pass = 1;
tile.quality = 50;
tile.sign[0][0] = SIGN_ZERO;
tile.sign[1][0] = SIGN_ZERO;
let coefficients = tile.coefficients;
let sign = tile.sign;
assert_eq!(
tile.decode_upgrade(
[&[0x90, 0x00], &[0x80, 0x00], &[]],
[&[], &[], &[]],
[ComponentCodecQuant::LOSSLESS; 3],
75,
),
Err(SrlError::Truncated)
);
assert_eq!(tile.coefficients, coefficients);
assert_eq!(tile.sign, sign);
assert_eq!(tile.prog_quant, [prev_prog_quant; 3]);
assert_eq!(tile.pass, 1);
assert_eq!(tile.quality, 50);
}
#[test]
fn tile_state_default_is_zeroed() {
let tile = TileState::new();
assert_eq!(tile.pass, 0);
assert_eq!(tile.quality, 0);
assert!(!tile.use_reduce_extrapolate);
assert!(tile.coefficients[0].iter().all(|&v| v == 0));
assert!(tile.sign[0].iter().all(|&v| v == 0));
}
#[test]
fn surface_tiles_dimensions() {
let surface = SurfaceTiles::new(1920, 1080, true).unwrap();
assert_eq!(surface.tiles_wide, 30);
assert_eq!(surface.tiles_high, 17);
assert!(surface.use_reduce_extrapolate);
}
#[test]
fn surface_tiles_exact_multiple() {
let surface = SurfaceTiles::new(1280, 768, false).unwrap();
assert_eq!(surface.tiles_wide, 20);
assert_eq!(surface.tiles_high, 12);
}
#[test]
fn surface_tiles_lazy_allocation() {
let mut surface = SurfaceTiles::new(128, 128, false).unwrap();
assert!(surface.get(0, 0).is_none());
let tile = surface.get_or_create(0, 0).unwrap();
assert_eq!(tile.pass, 0);
assert!(!tile.use_reduce_extrapolate);
assert!(surface.get(0, 0).is_some());
assert!(surface.get_or_create(2, 2).is_none());
}
#[test]
fn surface_tiles_reset() {
let mut surface = SurfaceTiles::new(128, 128, false).unwrap();
surface.get_or_create(0, 0);
assert!(surface.get(0, 0).is_some());
surface.reset();
assert!(surface.get(0, 0).is_none());
}
#[test]
fn decoder_new_is_empty() {
let decoder = ProgressiveDecoder::new();
assert!(decoder.contexts.is_empty());
}
#[test]
fn decoder_delete_nonexistent_context() {
let mut decoder = ProgressiveDecoder::new();
decoder.delete_context(1, 42);
}
#[test]
fn decoder_reset_clears_contexts_but_preserves_sub_band_references() {
let mut decoder = ProgressiveDecoder::new();
let result = decoder.decode_bitmap(1, 1, 64, 64, &simple_tile_stream(0, [64, -16, 24], true));
assert!(result.is_ok());
assert_eq!(decoder.contexts.len(), 1);
assert_eq!(decoder.references.len(), 1);
let reference = *decoder
.references
.get(&(1, 0, 0))
.expect("original tile reference should be retained");
decoder.reset();
assert!(decoder.contexts.is_empty());
assert_eq!(decoder.references.get(&(1, 0, 0)), Some(&reference));
assert!(
decoder
.decode_bitmap(
1,
2,
64,
64,
&simple_tile_stream(TILE_FLAG_DIFFERENCE, [7, -3, 5], true),
)
.is_ok(),
"a new codec context should use the retained surface reference"
);
}
#[test]
fn decoder_contexts_are_scoped_by_surface() {
let mut decoder = ProgressiveDecoder::new();
let stream = minimal_progressive_stream(true);
assert!(decoder.decode_bitmap(1, 0, 640, 480, &stream).is_ok());
assert!(decoder.decode_bitmap(2, 0, 800, 600, &stream).is_ok());
assert_eq!(decoder.contexts.len(), 2);
decoder.delete_context(1, 0);
assert_eq!(decoder.contexts.len(), 1);
assert!(decoder.contexts.contains_key(&(2, 0)));
assert!(decoder.decode_bitmap(1, 0, 640, 480, &stream).is_ok());
assert!(decoder.decode_bitmap(1, 1, 640, 480, &stream).is_ok());
assert_eq!(decoder.contexts.len(), 3);
decoder.delete_surface(1);
assert_eq!(decoder.contexts.len(), 1);
assert!(decoder.contexts.contains_key(&(2, 0)));
}
#[test]
fn decoder_context_fallback_is_scoped_by_surface() {
let mut decoder = ProgressiveDecoder::new();
let stream_with_context = minimal_progressive_stream(true);
let stream_without_context = minimal_progressive_stream(false);
assert!(decoder.decode_bitmap(1, 0, 640, 480, &stream_with_context).is_ok());
assert!(matches!(
decoder.decode_bitmap(2, 0, 640, 480, &stream_without_context),
Err(ProgressiveDecodeError::MissingBlock("CONTEXT"))
));
assert!(decoder.decode_bitmap(2, 0, 640, 480, &stream_with_context).is_ok());
assert!(decoder.decode_bitmap(2, 0, 640, 480, &stream_without_context).is_ok());
}
#[test]
fn decoder_error_display() {
let e = ProgressiveDecodeError::MissingBlock("SYNC");
assert!(e.to_string().contains("SYNC"));
let e = ProgressiveDecodeError::TileOutOfBounds { x_idx: 5, y_idx: 10 };
assert!(e.to_string().contains("5"));
assert!(e.to_string().contains("10"));
let e = ProgressiveDecodeError::InvalidQuantIndex { index: 3, table_len: 2 };
assert!(e.to_string().contains("3"));
}
#[test]
fn dequantize_component_ccq_shifts_correctly() {
let mut coefficients = vec![0i16; 4096];
coefficients[0] = 10; coefficients[4032] = 5;
let quant = ComponentCodecQuant {
ll3: 7,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 8,
lh1: 0,
hh1: 0,
};
dequantize_component_ccq(&mut coefficients, &quant, false);
assert_eq!(coefficients[0], 40);
assert_eq!(coefficients[4032], 10);
}
#[test]
fn rgba_to_ycbcr_pure_white() {
let pixels = vec![255u8; 64 * 64 * 4];
let mut y = vec![0i16; 4096];
let mut cb = vec![0i16; 4096];
let mut cr = vec![0i16; 4096];
rgba_to_ycbcr(&pixels, &mut y, &mut cb, &mut cr);
assert!((y[0] - 127).abs() <= 1, "Y for white: got {}", y[0]);
assert!(cb[0].abs() <= 1, "Cb for white: got {}", cb[0]);
assert!(cr[0].abs() <= 1, "Cr for white: got {}", cr[0]);
}
#[test]
fn rgba_to_ycbcr_pure_black() {
let pixels = vec![0u8; 64 * 64 * 4];
let mut y = vec![0i16; 4096];
let mut cb = vec![0i16; 4096];
let mut cr = vec![0i16; 4096];
rgba_to_ycbcr(&pixels, &mut y, &mut cb, &mut cr);
assert_eq!(y[0], -128);
assert_eq!(cb[0], 0);
assert_eq!(cr[0], 0);
}
#[test]
fn base_quantization_handles_all_wire_factors() {
for factor in 0..=15 {
let input = if factor < 6 { 1 } else { 1i16 << u32::from(factor - 6) };
let quant = ComponentCodecQuant {
ll3: factor,
hl3: factor,
lh3: factor,
hh3: factor,
hl2: factor,
lh2: factor,
hh2: factor,
hl1: factor,
lh1: factor,
hh1: factor,
};
let mut coefficients = [0i16; COEFFICIENTS_PER_COMPONENT];
coefficients[0] = input;
quantize_component_ccq(&mut coefficients, &quant, false);
dequantize_component_ccq(&mut coefficients, &quant, false);
assert_eq!(coefficients[0], input, "factor {factor}");
}
}
#[test]
fn base_dequantization_rounds_fractional_scales() {
let quant = ComponentCodecQuant {
ll3: 5,
hl3: 5,
lh3: 5,
hh3: 5,
hl2: 5,
lh2: 5,
hh2: 5,
hl1: 5,
lh1: 5,
hh1: 5,
};
let mut coefficients = [0i16; COEFFICIENTS_PER_COMPONENT];
coefficients[0] = 7;
coefficients[1] = -7;
dequantize_component_ccq(&mut coefficients, &quant, false);
assert_eq!(coefficients[0], 4);
assert_eq!(coefficients[1], -3);
}
#[test]
fn progressive_state_retains_base_quantized_coefficients() {
let base_quant = ComponentCodecQuant {
ll3: 5,
hl3: 5,
lh3: 5,
hh3: 5,
hl2: 5,
lh2: 5,
hh2: 5,
hl1: 5,
lh1: 5,
hh1: 5,
};
let prog_quant = ComponentCodecQuant {
ll3: 1,
hl3: 1,
lh3: 1,
hh3: 1,
hl2: 1,
lh2: 1,
hh2: 1,
hl1: 1,
lh1: 1,
hh1: 1,
};
let mut progressive_coefficients = [0i16; COEFFICIENTS_PER_COMPONENT];
progressive_coefficients[0] = 7;
progressive_coefficients[1] = -7;
let mut encoded = [0u8; 8192];
let encoded_len = crate::rlgr::encode(EntropyAlgorithm::Rlgr1, &progressive_coefficients, &mut encoded)
.expect("RLGR encoding should succeed");
let mut tile = TileState::new();
tile.decode_first(
[&encoded[..encoded_len]; 3],
[&base_quant; 3],
[prog_quant; 3],
[0; 3],
0,
false,
)
.expect("first-pass decoding should succeed");
assert_eq!(tile.coefficients[0][0], 14);
assert_eq!(tile.coefficients[0][1], -14);
let mut coefficients = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut sign = [SIGN_ZERO; COEFFICIENTS_PER_COMPONENT];
decode_first_pass(
&encoded[..encoded_len],
&base_quant,
&prog_quant,
false,
&mut coefficients,
&mut sign,
)
.expect("standalone first-pass decoding should succeed");
assert_eq!(coefficients[0], 7);
assert_eq!(coefficients[1], -7);
}
#[test]
#[expect(clippy::similar_names, reason = "Cb and Cr are standard YCbCr component names")]
fn progressive_fractional_base_quantization_reconstructs_rgb() {
let base_quant = ComponentCodecQuant {
ll3: 5,
hl3: 5,
lh3: 5,
hh3: 5,
hl2: 5,
lh2: 5,
hh2: 5,
hl1: 5,
lh1: 5,
hh1: 5,
};
let prog_quant = ComponentCodecQuant {
ll3: 1,
hl3: 1,
lh3: 1,
hh3: 1,
hl2: 1,
lh2: 1,
hh2: 1,
hl1: 1,
lh1: 1,
hh1: 1,
};
let expected = [64, 128, 192];
let mut pixels = vec![0u8; 64 * 64 * 4];
for pixel in pixels.chunks_exact_mut(4) {
pixel[..3].copy_from_slice(&expected);
pixel[3] = 0xFF;
}
let mut y = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut cb = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut cr = [0i16; COEFFICIENTS_PER_COMPONENT];
rgba_to_ycbcr(&pixels, &mut y, &mut cb, &mut cr);
let mut y_data = [0u8; 8192];
let mut cb_data = [0u8; 8192];
let mut cr_data = [0u8; 8192];
let y_len = encode_first_pass(&mut y, &mut y_data, &base_quant, &prog_quant, false)
.expect("Y first-pass encoding should succeed");
let cb_len = encode_first_pass(&mut cb, &mut cb_data, &base_quant, &prog_quant, false)
.expect("Cb first-pass encoding should succeed");
let cr_len = encode_first_pass(&mut cr, &mut cr_data, &base_quant, &prog_quant, false)
.expect("Cr first-pass encoding should succeed");
let mut tile = TileState::new();
tile.decode_first(
[&y_data[..y_len], &cb_data[..cb_len], &cr_data[..cr_len]],
[&base_quant; 3],
[prog_quant; 3],
[0; 3],
0,
false,
)
.expect("first-pass decoding should succeed");
let mut actual = vec![0u8; 64 * 64 * 4];
tile.reconstruct_to_rgba(&mut actual);
for actual in actual.chunks_exact(4) {
for channel in 0..3 {
let difference = i16::from(expected[channel]) - i16::from(actual[channel]);
assert!(difference.abs() <= 2, "expected {expected:?}, got {:?}", &actual[..3]);
}
assert_eq!(actual[3], 0xFF);
}
}
#[test]
fn progressive_q6_reconstructs_rgb_color_vectors() {
let base_quant = ComponentCodecQuant {
ll3: 6,
hl3: 6,
lh3: 6,
hh3: 6,
hl2: 6,
lh2: 6,
hh2: 6,
hl1: 6,
lh1: 6,
hh1: 6,
};
for expected in [
[0, 0, 0],
[255, 255, 255],
[255, 0, 0],
[0, 255, 0],
[0, 0, 255],
[64, 128, 192],
] {
let mut pixels = vec![0u8; 64 * 64 * 4];
for pixel in pixels.chunks_exact_mut(4) {
pixel[..3].copy_from_slice(&expected);
pixel[3] = 0xFF;
}
let mut y = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut cb = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut cr = [0i16; COEFFICIENTS_PER_COMPONENT];
rgba_to_ycbcr(&pixels, &mut y, &mut cb, &mut cr);
let mut temp = [0i16; COEFFICIENTS_PER_COMPONENT];
crate::dwt::encode(&mut y, &mut temp);
crate::dwt::encode(&mut cb, &mut temp);
crate::dwt::encode(&mut cr, &mut temp);
quantize_component_ccq(&mut y, &base_quant, false);
quantize_component_ccq(&mut cb, &base_quant, false);
quantize_component_ccq(&mut cr, &base_quant, false);
dequantize_component_ccq(&mut y, &base_quant, false);
dequantize_component_ccq(&mut cb, &base_quant, false);
dequantize_component_ccq(&mut cr, &base_quant, false);
let mut tile = TileState::new();
tile.coefficients = [y, cb, cr];
let mut actual = vec![0u8; 64 * 64 * 4];
tile.reconstruct_to_rgba(&mut actual);
for actual in actual.chunks_exact(4) {
for channel in 0..3 {
let difference = i16::from(expected[channel]) - i16::from(actual[channel]);
assert!(difference.abs() <= 2, "expected {:?}, got {:?}", expected, &actual[..3]);
}
assert_eq!(actual[3], 0xFF);
}
}
}
#[test]
fn quantize_ccq_scales_coefficients() {
let mut coefficients = [0i16; 4096];
coefficients[0] = 40; coefficients[4032] = 10;
let quant = ComponentCodecQuant {
ll3: 7,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 8,
lh1: 0,
hh1: 0,
};
quantize_component_ccq(&mut coefficients, &quant, false);
assert_eq!(coefficients[0], 10);
assert_eq!(coefficients[4032], 5);
}
#[test]
fn quantize_ccq_preserves_negative_sign() {
let mut coefficients = [0i16; 4096];
coefficients[0] = -40;
let quant = ComponentCodecQuant {
ll3: 0,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 8,
lh1: 0,
hh1: 0,
};
quantize_component_ccq(&mut coefficients, &quant, false);
assert_eq!(coefficients[0], -10);
}
#[test]
fn raw_bit_writer_single_byte() {
let mut w = RawBitWriter::new();
w.write_bits(0xA5, 8);
assert_eq!(w.finish(), vec![0xA5]);
}
#[test]
fn raw_bit_writer_partial_byte_padded() {
let mut w = RawBitWriter::new();
w.write_bits(0b101, 3);
assert_eq!(w.finish(), vec![0xA0]);
}
#[test]
fn raw_bit_writer_multi_byte() {
let mut w = RawBitWriter::new();
w.write_bits(0xFF, 8);
w.write_bits(0b1010, 4);
assert_eq!(w.finish(), vec![0xFF, 0xA0]);
}
#[test]
fn encode_first_pass_produces_output() {
let mut coefficients = [100i16; 4096];
let mut output = vec![0u8; 8192];
let base_quant = ComponentCodecQuant::LOSSLESS;
let prog_quant = ComponentCodecQuant::LOSSLESS;
let result = encode_first_pass(&mut coefficients, &mut output, &base_quant, &prog_quant, false);
assert!(result.is_ok(), "RLGR encode failed: {:?}", result.err());
let bytes_written = result.unwrap();
assert!(bytes_written > 0, "expected non-zero encoded output");
assert!(bytes_written < 8192, "flat tile should compress");
}
#[test]
fn encode_first_pass_reduce_extrapolate() {
let mut coefficients = [50i16; 4096];
let mut output = vec![0u8; 8192];
let base_quant = ComponentCodecQuant::LOSSLESS;
let prog_quant = ComponentCodecQuant::LOSSLESS;
let result = encode_first_pass(
&mut coefficients,
&mut output,
&base_quant,
&prog_quant,
true, );
assert!(result.is_ok(), "RLGR encode failed: {:?}", result.err());
assert!(result.unwrap() > 0);
}
#[test]
fn encode_upgrade_pass_empty_when_no_refinement() {
let coefficients = [0i16; 4096];
let prev_coefficients = [0i16; 4096];
let sign = [SIGN_ZERO; 4096];
let prog_quant = ComponentCodecQuant::LOSSLESS;
let (srl_data, raw_data) = encode_upgrade_pass(
&coefficients,
&prev_coefficients,
&prog_quant,
&prog_quant,
&sign,
false,
)
.unwrap();
assert!(srl_data.is_empty(), "no refinement bits, SRL should be empty");
assert!(raw_data.is_empty(), "no refinement bits, raw should be empty");
}
#[test]
fn first_pass_encode_decode_round_trip_lossless() {
let original = [42i16; COEFFICIENTS_PER_COMPONENT];
let mut encode_buf = original;
let mut output = vec![0u8; 16384];
let base_quant = ComponentCodecQuant::LOSSLESS;
let prog_quant = ComponentCodecQuant::LOSSLESS;
let bytes = encode_first_pass(&mut encode_buf, &mut output, &base_quant, &prog_quant, false).unwrap();
let mut decoded = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut sign = [0i8; COEFFICIENTS_PER_COMPONENT];
decode_first_pass(
&output[..bytes],
&base_quant,
&prog_quant,
false,
&mut decoded,
&mut sign,
)
.unwrap();
let mut temp = [0i16; COEFFICIENTS_PER_COMPONENT];
crate::dwt::decode(&mut decoded, &mut temp);
let max_err = original
.iter()
.zip(decoded.iter())
.map(|(a, b)| (i32::from(*a) - i32::from(*b)).unsigned_abs())
.max()
.unwrap();
assert!(max_err <= 4, "flat data round-trip max error {max_err} exceeds 4");
}
#[test]
fn first_pass_encode_decode_round_trip_reduce_extrapolate() {
let original = [42i16; COEFFICIENTS_PER_COMPONENT];
let mut encode_buf = original;
let mut output = vec![0u8; 16384];
let base_quant = ComponentCodecQuant::LOSSLESS;
let prog_quant = ComponentCodecQuant::LOSSLESS;
let bytes = encode_first_pass(&mut encode_buf, &mut output, &base_quant, &prog_quant, true).unwrap();
let mut decoded = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut sign = [0i8; COEFFICIENTS_PER_COMPONENT];
decode_first_pass(
&output[..bytes],
&base_quant,
&prog_quant,
true,
&mut decoded,
&mut sign,
)
.unwrap();
let mut temp = [0i16; COEFFICIENTS_PER_COMPONENT];
crate::dwt_extrapolate::decode(&mut decoded, &mut temp);
let max_err = original
.iter()
.zip(decoded.iter())
.map(|(a, b)| (i32::from(*a) - i32::from(*b)).unsigned_abs())
.max()
.unwrap();
assert!(
max_err <= 6,
"reduce-extrapolate round-trip max error {max_err} exceeds 6"
);
}
#[test]
fn first_pass_encode_decode_with_quantization() {
let mut coefficients = [42i16; COEFFICIENTS_PER_COMPONENT];
let mut output = vec![0u8; 16384];
let base_quant = ComponentCodecQuant {
ll3: 6,
hl3: 6,
lh3: 6,
hh3: 6,
hl2: 7,
lh2: 7,
hh2: 7,
hl1: 8,
lh1: 8,
hh1: 8,
};
let prog_quant = ComponentCodecQuant::LOSSLESS;
let bytes = encode_first_pass(&mut coefficients, &mut output, &base_quant, &prog_quant, false).unwrap();
assert!(bytes > 0, "should produce encoded output");
let mut decoded = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut sign = [0i8; COEFFICIENTS_PER_COMPONENT];
decode_first_pass(
&output[..bytes],
&base_quant,
&prog_quant,
false,
&mut decoded,
&mut sign,
)
.unwrap();
let mut temp = [0i16; COEFFICIENTS_PER_COMPONENT];
crate::dwt::decode(&mut decoded, &mut temp);
let mean_err: f64 = decoded
.iter()
.map(|v| f64::from((i32::from(*v) - 42).unsigned_abs()))
.sum::<f64>()
/ 4096.0;
assert!(
mean_err < 200.0,
"mean error {mean_err} too large for quantized flat tile"
);
}
#[test]
#[expect(clippy::similar_names, reason = "y/cb/cr are standard YCbCr component names")]
fn rgba_ycbcr_reconstruct_round_trip() {
let mut pixels = vec![0u8; 64 * 64 * 4];
for i in 0..64 * 64 {
let row = i / 64;
let col = i % 64;
pixels[i * 4] = (row * 4) as u8; pixels[i * 4 + 1] = (col * 4) as u8; pixels[i * 4 + 2] = 128; pixels[i * 4 + 3] = 255; }
let mut y = vec![0i16; 4096];
let mut cb = vec![0i16; 4096];
let mut cr = vec![0i16; 4096];
rgba_to_ycbcr(&pixels, &mut y, &mut cb, &mut cr);
for i in 0..4096 {
assert!(y[i] >= -128 && y[i] <= 127, "Y[{i}] = {} out of range", y[i]);
assert!(cb[i] >= -128 && cb[i] <= 127, "Cb[{i}] = {} out of range", cb[i]);
assert!(cr[i] >= -128 && cr[i] <= 127, "Cr[{i}] = {} out of range", cr[i]);
}
let mut max_err = 0i32;
for i in 0..64 * 64 {
let y_val = i32::from(y[i]) + 128;
let cb_val = i32::from(cb[i]);
let cr_val = i32::from(cr[i]);
let r_rec = (y_val + ((cr_val * 91881 + 32768) >> 16)).clamp(0, 255);
let g_rec = (y_val - ((cb_val * 22554 + cr_val * 46802 + 32768) >> 16)).clamp(0, 255);
let b_rec = (y_val + ((cb_val * 116130 + 32768) >> 16)).clamp(0, 255);
let off = i * 4;
let r_orig = i32::from(pixels[off]);
let g_orig = i32::from(pixels[off + 1]);
let b_orig = i32::from(pixels[off + 2]);
max_err = max_err.max((r_rec - r_orig).abs());
max_err = max_err.max((g_rec - g_orig).abs());
max_err = max_err.max((b_rec - b_orig).abs());
}
assert!(
max_err <= 2,
"RGB -> YCbCr -> RGB max per-channel error {max_err} exceeds 2"
);
}
#[test]
fn upgrade_pass_encode_decode_round_trip() {
let prev_prog_quant = ComponentCodecQuant {
ll3: 0,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 4,
lh1: 0,
hh1: 0,
};
let curr_prog_quant = ComponentCodecQuant {
ll3: 0,
hl3: 0,
lh3: 0,
hh3: 0,
hl2: 0,
lh2: 0,
hh2: 0,
hl1: 2,
lh1: 0,
hh1: 0,
};
let mut prev_coeffs = vec![0i16; COEFFICIENTS_PER_COMPONENT];
let mut refined_coeffs = vec![0i16; COEFFICIENTS_PER_COMPONENT];
let mut sign = vec![SIGN_ZERO; COEFFICIENTS_PER_COMPONENT];
for i in 0..1024 {
let base = ((i as i32) * 5) % 256 - 128;
let coarse = base & !0x0F;
let refined = base;
prev_coeffs[i] = coarse as i16;
refined_coeffs[i] = refined as i16;
sign[i] = match prev_coeffs[i].cmp(&0) {
core::cmp::Ordering::Greater => SIGN_POSITIVE,
core::cmp::Ordering::Less => SIGN_NEGATIVE,
core::cmp::Ordering::Equal => SIGN_ZERO,
};
}
let prev_dist: u32 = prev_coeffs
.iter()
.zip(refined_coeffs.iter())
.map(|(p, r)| (i32::from(*p) - i32::from(*r)).unsigned_abs())
.sum();
let (srl_data, raw_data) = encode_upgrade_pass(
&refined_coeffs,
&prev_coeffs,
&prev_prog_quant,
&curr_prog_quant,
&sign,
false,
)
.unwrap();
let mut decoded = prev_coeffs.clone();
let mut decoded_sign = sign.clone();
decode_upgrade_pass(
&srl_data,
&raw_data,
&prev_prog_quant,
&curr_prog_quant,
false,
&mut decoded,
&mut decoded_sign,
)
.unwrap();
let post_dist: u32 = decoded
.iter()
.zip(refined_coeffs.iter())
.map(|(d, r)| (i32::from(*d) - i32::from(*r)).unsigned_abs())
.sum();
assert!(
post_dist <= prev_dist,
"upgrade pass must not increase distance to refined: prev_dist={prev_dist} post_dist={post_dist}"
);
}
#[test]
fn quantize_dequantize_ccq_round_trip() {
let quant = ComponentCodecQuant {
ll3: 4,
hl3: 4,
lh3: 4,
hh3: 5,
hl2: 5,
lh2: 5,
hh2: 6,
hl1: 6,
lh1: 6,
hh1: 7,
};
let original = {
let mut c = [0i16; COEFFICIENTS_PER_COMPONENT];
for (i, v) in c.iter_mut().enumerate() {
*v = ((i * 7 % 256) as i16) - 128;
}
c
};
let mut coefficients = original;
quantize_component_ccq(&mut coefficients, &quant, false);
dequantize_component_ccq(&mut coefficients, &quant, false);
let max_err = original
.iter()
.zip(coefficients.iter())
.map(|(a, b)| (i32::from(*a) - i32::from(*b)).unsigned_abs())
.max()
.unwrap();
assert!(
max_err <= 64,
"quantize/dequantize round-trip max error {max_err} exceeds 64"
);
}
fn encode_full_quality_component(value: i16) -> Vec<u8> {
let mut coefficients = [value; COEFFICIENTS_PER_COMPONENT];
let mut encoded = vec![0; 8192];
let encoded_len = encode_first_pass(
&mut coefficients,
&mut encoded,
&ComponentCodecQuant::LOSSLESS,
&ComponentCodecQuant::LOSSLESS,
false,
)
.expect("full-quality component encoding should succeed");
encoded.truncate(encoded_len);
encoded
}
fn progressive_tile_stream(
include_context: bool,
quant_vals: Vec<ComponentCodecQuant>,
quant_prog_vals: Vec<rat_rdp_pdu::codecs::rfx::progressive::ProgressiveCodecQuant>,
tile: rat_rdp_pdu::codecs::rfx::progressive::ProgressiveTile<'_>,
) -> Vec<u8> {
use rat_rdp_pdu::codecs::rfx::RfxRectangle;
use rat_rdp_pdu::codecs::rfx::progressive::{
ProgressiveBlock, ProgressiveContextPdu, ProgressiveFrameBeginPdu, ProgressiveFrameEndPdu,
ProgressiveRegion, ProgressiveSyncPdu, encode_progressive_stream,
};
let mut blocks = vec![ProgressiveBlock::Sync(ProgressiveSyncPdu)];
if include_context {
blocks.push(ProgressiveBlock::Context(ProgressiveContextPdu {
context_id: 0,
tile_size: 0x0040,
flags: 0,
}));
}
blocks.extend([
ProgressiveBlock::FrameBegin(ProgressiveFrameBeginPdu {
frame_index: 0,
region_count: 1,
}),
ProgressiveBlock::Region(ProgressiveRegion {
tile_size: 0x40,
rects: vec![RfxRectangle {
x: 0,
y: 0,
width: 64,
height: 64,
}],
quant_vals,
quant_prog_vals,
flags: 0,
tiles: vec![tile],
}),
ProgressiveBlock::FrameEnd(ProgressiveFrameEndPdu),
]);
encode_progressive_stream(&blocks).expect("synthetic progressive stream should encode")
}
fn simple_tile_stream(flags: u8, components: [i16; 3], include_context: bool) -> Vec<u8> {
use rat_rdp_pdu::codecs::rfx::progressive::{ProgressiveTile, TileSimple};
let component_data = components.map(encode_full_quality_component);
progressive_tile_stream(
include_context,
vec![ComponentCodecQuant::LOSSLESS],
vec![],
ProgressiveTile::Simple(TileSimple {
quant_idx_y: 0,
quant_idx_cb: 0,
quant_idx_cr: 0,
x_idx: 0,
y_idx: 0,
flags,
y_data: &component_data[0],
cb_data: &component_data[1],
cr_data: &component_data[2],
tail_data: &[],
}),
)
}
fn first_tile_stream(
flags: u8,
component_data: &[u8],
base_quant: ComponentCodecQuant,
progressive_quant: ComponentCodecQuant,
include_context: bool,
) -> Vec<u8> {
use rat_rdp_pdu::codecs::rfx::progressive::{ProgressiveCodecQuant, ProgressiveTile, TileFirst};
progressive_tile_stream(
include_context,
vec![base_quant],
vec![ProgressiveCodecQuant {
quality: 0,
y_quant: progressive_quant,
cb_quant: progressive_quant,
cr_quant: progressive_quant,
}],
ProgressiveTile::First(TileFirst {
quant_idx_y: 0,
quant_idx_cb: 0,
quant_idx_cr: 0,
x_idx: 0,
y_idx: 0,
flags,
quality: 0,
y_data: component_data,
cb_data: component_data,
cr_data: component_data,
tail_data: &[],
}),
)
}
fn upgrade_tile_stream(
raw_data: &[u8],
base_quant: ComponentCodecQuant,
first_progressive_quant: ComponentCodecQuant,
upgrade_progressive_quant: ComponentCodecQuant,
include_context: bool,
) -> Vec<u8> {
use rat_rdp_pdu::codecs::rfx::progressive::{ProgressiveCodecQuant, ProgressiveTile, TileUpgrade};
progressive_tile_stream(
include_context,
vec![base_quant],
vec![
ProgressiveCodecQuant {
quality: 0,
y_quant: first_progressive_quant,
cb_quant: first_progressive_quant,
cr_quant: first_progressive_quant,
},
ProgressiveCodecQuant {
quality: 1,
y_quant: upgrade_progressive_quant,
cb_quant: upgrade_progressive_quant,
cr_quant: upgrade_progressive_quant,
},
],
ProgressiveTile::Upgrade(TileUpgrade {
quant_idx_y: 0,
quant_idx_cb: 0,
quant_idx_cr: 0,
x_idx: 0,
y_idx: 0,
quality: 1,
y_srl_data: &[],
y_raw_data: raw_data,
cb_srl_data: &[],
cb_raw_data: raw_data,
cr_srl_data: &[],
cr_raw_data: raw_data,
}),
)
}
fn decode_full_quality_components(components: [i16; 3]) -> DecDwtQ {
let component_data = components.map(encode_full_quality_component);
let mut state = TileState::new();
state
.decode_first(
[&component_data[0], &component_data[1], &component_data[2]],
[&ComponentCodecQuant::LOSSLESS; 3],
[ComponentCodecQuant::LOSSLESS; 3],
[0; 3],
0xFF,
false,
)
.expect("full-quality components should decode");
state.coefficients
}
#[test]
fn difference_tile_adds_to_its_retained_surface_reference() {
let original_components = [64, -16, 24];
let other_surface_components = [-48, 8, 40];
let difference_components = [7, -3, 5];
let mut decoder = ProgressiveDecoder::new();
let first_pixels = decoder
.decode_bitmap(1, 7, 64, 64, &simple_tile_stream(0, original_components, true))
.expect("original tile should decode")
.pop()
.expect("original tile should produce an update")
.pixels;
let reference = decoder
.contexts
.get(&(1, 7))
.and_then(|context| context.surface.get(0, 0))
.expect("original tile state should be retained")
.coefficients;
decoder
.decode_bitmap(2, 8, 64, 64, &simple_tile_stream(0, other_surface_components, true))
.expect("other surface tile should decode");
let other_reference = decoder
.contexts
.get(&(2, 8))
.and_then(|context| context.surface.get(0, 0))
.expect("other surface tile state should be retained")
.coefficients;
let expected_delta = decode_full_quality_components(difference_components);
let difference_pixels = decoder
.decode_bitmap(
1,
7,
64,
64,
&simple_tile_stream(TILE_FLAG_DIFFERENCE, difference_components, false),
)
.expect("difference tile should decode")
.pop()
.expect("difference tile should produce an update")
.pixels;
assert_ne!(first_pixels, difference_pixels);
let updated_tile = decoder
.contexts
.get(&(1, 7))
.and_then(|context| context.surface.get(0, 0))
.expect("difference tile state should be retained");
assert!(updated_tile.is_difference);
for ((updated_component, reference_component), delta_component) in updated_tile
.coefficients
.iter()
.zip(reference.iter())
.zip(expected_delta.iter())
{
for ((updated, retained), delta) in updated_component
.iter()
.zip(reference_component.iter())
.zip(delta_component.iter())
{
assert_eq!(*updated, retained.saturating_add(*delta));
}
}
assert_eq!(
decoder
.contexts
.get(&(2, 8))
.and_then(|context| context.surface.get(0, 0))
.expect("other surface tile state should remain retained")
.coefficients,
other_reference
);
}
#[test]
fn first_difference_tile_adds_to_its_retained_surface_reference() {
let progressive_quant = ComponentCodecQuant::LOSSLESS;
let original_data = encode_full_quality_component(64);
let difference_data = encode_full_quality_component(7);
let summed_data = encode_full_quality_component(71);
let mut decoder = ProgressiveDecoder::new();
decoder
.decode_bitmap(
1,
7,
64,
64,
&first_tile_stream(
0,
&original_data,
ComponentCodecQuant::LOSSLESS,
progressive_quant,
true,
),
)
.expect("original first-pass tile should decode");
let accumulated = decoder
.decode_bitmap(
1,
7,
64,
64,
&first_tile_stream(
TILE_FLAG_DIFFERENCE,
&difference_data,
ComponentCodecQuant::LOSSLESS,
progressive_quant,
false,
),
)
.expect("difference first-pass tile should decode");
let mut summed_decoder = ProgressiveDecoder::new();
let summed = summed_decoder
.decode_bitmap(
1,
7,
64,
64,
&first_tile_stream(0, &summed_data, ComponentCodecQuant::LOSSLESS, progressive_quant, true),
)
.expect("summed first-pass tile should decode");
assert_eq!(accumulated[0].pixels, summed[0].pixels);
assert_eq!(
decoder.references.get(&(1, 0, 0)),
summed_decoder.references.get(&(1, 0, 0))
);
assert!(
decoder
.contexts
.get(&(1, 7))
.and_then(|context| context.surface.get(0, 0))
.expect("difference tile state should be retained")
.is_difference
);
}
#[test]
fn original_tile_replaces_a_retained_surface_reference() {
let original_components = [64, -16, 24];
let difference_components = [7, -3, 5];
let replacement_components = [-48, 8, 40];
let mut decoder = ProgressiveDecoder::new();
decoder
.decode_bitmap(1, 7, 64, 64, &simple_tile_stream(0, original_components, true))
.expect("original tile should decode");
decoder
.decode_bitmap(
1,
7,
64,
64,
&simple_tile_stream(TILE_FLAG_DIFFERENCE, difference_components, false),
)
.expect("difference tile should decode");
let replacement = decoder
.decode_bitmap(1, 7, 64, 64, &simple_tile_stream(0, replacement_components, false))
.expect("replacement tile should decode");
let mut standalone_decoder = ProgressiveDecoder::new();
let standalone = standalone_decoder
.decode_bitmap(1, 7, 64, 64, &simple_tile_stream(0, replacement_components, true))
.expect("standalone replacement tile should decode");
assert_eq!(replacement[0].pixels, standalone[0].pixels);
assert_eq!(
decoder.references.get(&(1, 0, 0)),
standalone_decoder.references.get(&(1, 0, 0))
);
assert!(
!decoder
.contexts
.get(&(1, 7))
.and_then(|context| context.surface.get(0, 0))
.expect("replacement tile state should be retained")
.is_difference
);
}
#[test]
fn difference_tile_uses_the_reference_updated_by_an_upgrade() {
let base_quant = ComponentCodecQuant {
ll3: 6,
hl3: 6,
lh3: 6,
hh3: 6,
hl2: 6,
lh2: 6,
hh2: 6,
hl1: 6,
lh1: 6,
hh1: 6,
};
let mut progressive_quant = ComponentCodecQuant::LOSSLESS;
progressive_quant.ll3 = 1;
let mut original_coefficients = [0; COEFFICIENTS_PER_COMPONENT];
original_coefficients[4032] = 25;
let mut difference_coefficients = [0; COEFFICIENTS_PER_COMPONENT];
difference_coefficients[4032] = 5;
let original_data = {
let mut encoded = vec![0; 16 * 1024];
let len = crate::rlgr::encode(EntropyAlgorithm::Rlgr1, &original_coefficients, &mut encoded)
.expect("original RLGR encoding should succeed");
encoded.truncate(len);
encoded
};
let difference_data = {
let mut encoded = vec![0; 16 * 1024];
let len = crate::rlgr::encode(EntropyAlgorithm::Rlgr1, &difference_coefficients, &mut encoded)
.expect("difference RLGR encoding should succeed");
encoded.truncate(len);
encoded
};
let mut decoder = ProgressiveDecoder::new();
decoder
.decode_bitmap(
1,
7,
64,
64,
&first_tile_stream(0, &original_data, base_quant, progressive_quant, true),
)
.expect("original first-pass tile should decode");
let initial_reference = *decoder
.references
.get(&(1, 0, 0))
.expect("original tile should retain a reference");
let raw_data = [0xFF; 8];
decoder
.decode_bitmap(
1,
7,
64,
64,
&upgrade_tile_stream(
&raw_data,
base_quant,
progressive_quant,
ComponentCodecQuant::LOSSLESS,
false,
),
)
.expect("upgrade tile should decode");
let upgraded_reference = *decoder
.references
.get(&(1, 0, 0))
.expect("upgrade should update the retained reference");
assert_ne!(upgraded_reference[0][4032], initial_reference[0][4032]);
decoder
.decode_bitmap(
1,
7,
64,
64,
&first_tile_stream(
TILE_FLAG_DIFFERENCE,
&difference_data,
base_quant,
progressive_quant,
false,
),
)
.expect("difference first-pass tile should decode");
let updated_reference = decoder
.references
.get(&(1, 0, 0))
.expect("difference tile should update the retained reference");
let mut delta = TileState::new();
delta
.decode_first(
[&difference_data; 3],
[&base_quant; 3],
[progressive_quant; 3],
[0; 3],
0,
false,
)
.expect("difference tile payload should decode independently");
for ((updated_component, upgraded_component), delta_component) in updated_reference
.iter()
.zip(upgraded_reference.iter())
.zip(delta.coefficients.iter())
{
for ((updated, upgraded), delta) in updated_component
.iter()
.zip(upgraded_component.iter())
.zip(delta_component.iter())
{
assert_eq!(*updated, upgraded.saturating_add(*delta));
}
}
}
#[test]
fn difference_tile_reference_survives_codec_context_deletion() {
let original_components = [64, -16, 24];
let difference_components = [7, -3, 5];
let mut decoder = ProgressiveDecoder::new();
decoder
.decode_bitmap(1, 7, 64, 64, &simple_tile_stream(0, original_components, true))
.expect("original tile should decode");
let reference = *decoder
.references
.get(&(1, 0, 0))
.expect("original tile reference should be retained");
decoder.delete_context(1, 7);
assert!(decoder.references.contains_key(&(1, 0, 0)));
decoder
.decode_bitmap(
1,
8,
64,
64,
&simple_tile_stream(TILE_FLAG_DIFFERENCE, difference_components, true),
)
.expect("difference tile should decode with the surface reference");
let expected_delta = decode_full_quality_components(difference_components);
let updated = decoder
.contexts
.get(&(1, 8))
.and_then(|context| context.surface.get(0, 0))
.expect("difference tile state should be retained")
.coefficients;
for ((updated_component, reference_component), delta_component) in
updated.iter().zip(reference.iter()).zip(expected_delta.iter())
{
for ((updated, retained), delta) in updated_component
.iter()
.zip(reference_component.iter())
.zip(delta_component.iter())
{
assert_eq!(*updated, retained.saturating_add(*delta));
}
}
assert_eq!(decoder.references.get(&(1, 0, 0)), Some(&updated));
decoder.delete_surface(1);
assert!(!decoder.references.contains_key(&(1, 0, 0)));
assert!(matches!(
decoder.decode_bitmap(
1,
9,
64,
64,
&simple_tile_stream(TILE_FLAG_DIFFERENCE, difference_components, true),
),
Err(ProgressiveDecodeError::MissingTileReference { x_idx: 0, y_idx: 0 })
));
}
#[test]
fn resizing_a_surface_preserves_sub_band_references() {
let mut decoder = ProgressiveDecoder::new();
let original_components = [64, -16, 24];
let difference_components = [7, -3, 5];
decoder
.decode_bitmap(1, 7, 64, 64, &simple_tile_stream(0, original_components, true))
.expect("original tile should decode");
let reference = *decoder
.references
.get(&(1, 0, 0))
.expect("original tile reference should be retained");
decoder
.decode_bitmap(1, 7, 128, 64, &minimal_progressive_stream(true))
.expect("resized surface should decode");
assert_eq!(decoder.references.get(&(1, 0, 0)), Some(&reference));
decoder
.decode_bitmap(
1,
7,
128,
64,
&simple_tile_stream(TILE_FLAG_DIFFERENCE, difference_components, false),
)
.expect("difference tile should decode after resizing");
let expected_delta = decode_full_quality_components(difference_components);
let updated = decoder
.references
.get(&(1, 0, 0))
.expect("difference tile should update the retained reference");
for ((updated_component, reference_component), delta_component) in
updated.iter().zip(reference.iter()).zip(expected_delta.iter())
{
for ((updated_coefficient, reference_coefficient), delta_coefficient) in updated_component
.iter()
.zip(reference_component.iter())
.zip(delta_component.iter())
{
assert_eq!(
*updated_coefficient,
reference_coefficient.saturating_add(*delta_coefficient)
);
}
}
}
#[test]
fn difference_tile_requires_a_retained_reference() {
let mut decoder = ProgressiveDecoder::new();
assert!(matches!(
decoder.decode_bitmap(
1,
7,
64,
64,
&simple_tile_stream(TILE_FLAG_DIFFERENCE, [7, -3, 5], true),
),
Err(ProgressiveDecodeError::MissingTileReference { x_idx: 0, y_idx: 0 })
));
}
#[test]
fn failed_difference_tile_keeps_retained_state() {
let original_data = [64, -16, 24].map(encode_full_quality_component);
let difference_data = [7, -3, 5].map(encode_full_quality_component);
let mut state = TileState::new();
let base_quants = [&ComponentCodecQuant::LOSSLESS; 3];
let prog_quants = [ComponentCodecQuant::LOSSLESS; 3];
state
.decode_first(
[&original_data[0], &original_data[1], &original_data[2]],
base_quants,
prog_quants,
[0; 3],
0xFF,
false,
)
.expect("original tile should decode");
let coefficients = state.coefficients;
let sign = state.sign;
let prog_quant = state.prog_quant;
let quant_idx = state.quant_idx;
let base_quant = state.base_quant;
let pass = state.pass;
let is_difference = state.is_difference;
let quality = state.quality;
let use_reduce_extrapolate = state.use_reduce_extrapolate;
assert!(
state
.decode_first_with_difference(
[&difference_data[0], &difference_data[1], &[]],
base_quants,
prog_quants,
Some(&coefficients),
FirstPassOptions {
quant_idx: [0; 3],
quality: 0xFF,
use_reduce_extrapolate: false,
},
)
.is_err(),
"invalid difference payload should fail"
);
assert_eq!(state.coefficients, coefficients);
assert_eq!(state.sign, sign);
assert_eq!(state.prog_quant, prog_quant);
assert_eq!(state.quant_idx, quant_idx);
assert_eq!(state.base_quant, base_quant);
assert_eq!(state.pass, pass);
assert_eq!(state.is_difference, is_difference);
assert_eq!(state.quality, quality);
assert_eq!(state.use_reduce_extrapolate, use_reduce_extrapolate);
let mut delta = TileState::new();
delta
.decode_first(
[&difference_data[0], &difference_data[1], &difference_data[2]],
base_quants,
prog_quants,
[0; 3],
0xFF,
false,
)
.expect("valid difference payload should decode");
state
.decode_first_with_difference(
[&difference_data[0], &difference_data[1], &difference_data[2]],
base_quants,
prog_quants,
Some(&coefficients),
FirstPassOptions {
quant_idx: [0; 3],
quality: 0xFF,
use_reduce_extrapolate: false,
},
)
.expect("difference tile after failed decode should use retained state");
for ((updated_component, retained_component), delta_component) in state
.coefficients
.iter()
.zip(coefficients.iter())
.zip(delta.coefficients.iter())
{
for ((updated, retained), delta) in updated_component
.iter()
.zip(retained_component.iter())
.zip(delta_component.iter())
{
assert_eq!(*updated, retained.saturating_add(*delta));
}
}
}
}