1//! Decompression algorithm from:
2//! <https://problemkaputt.github.io/fullsnes.htm#snescartsdd1decompressionalgorithm>
3//!
4//! The algorithm is also described in English here:
5//! <https://wiki.superfamicom.org/s-dd1>
6
7use crate::sdd1::Sdd1Mmc;
8use bincode::{Decode, Encode};
9use jgenesis_common::num::GetBit;
10
11// Golomb decoder codeword size, indexed by state
12// Higher states use longer codewords because runs are theoretically more likely to end with the MPS
13// States 25-32 don't follow the pattern because they're highly adaptable states that are only used
14// shortly after initialization (see MPS/LPS evolution tables)
15#[rustfmt::skip]
16const EVOLUTION_CODE_SIZE: &[u8; 33] = &[
17    0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3,
18    4, 4, 5, 5, 6, 6, 7, 7, 0, 1, 2, 3, 4, 5, 6, 7
19];
20
21// MPS = Most probable symbol
22// If a run ends in the MPS, move to a higher state
23#[rustfmt::skip]
24const EVOLUTION_MPS_NEXT: &[u8; 33] = &[
25    25, 2, 3, 4, 5, 6, 7, 8, 9,10,11,12,13,14,15,16,17,
26    18,19,20,21,22,23,24,24,26,27,28,29,30,31,32,24
27];
28
29// LPS = Least probable symbol
30// If a run ends in the LPS, move to a lower state
31#[rustfmt::skip]
32const EVOLUTION_LPS_NEXT: &[u8; 33] = &[
33    25, 1, 1, 2, 3, 4, 5, 6, 7, 8, 9,10,11,12,13,14,15,
34    16,17,18,19,20,21,22,23, 1, 2, 4, 8,12,16,18,22
35];
36
37#[rustfmt::skip]
38const RUN_TABLE: &[u8; 128] = &[
39    128, 64, 96, 32, 112, 48, 80, 16, 120, 56, 88, 24, 104, 40, 72, 8,
40    124, 60, 92, 28, 108, 44, 76, 12, 116, 52, 84, 20, 100, 36, 68, 4,
41    126, 62, 94, 30, 110, 46, 78, 14, 118, 54, 86, 22, 102, 38, 70, 6,
42    122, 58, 90, 26, 106, 42, 74, 10, 114, 50, 82, 18,  98, 34, 66, 2,
43    127, 63, 95, 31, 111, 47, 79, 15, 119, 55, 87, 23, 103, 39, 71, 7,
44    123, 59, 91, 27, 107, 43, 75, 11, 115, 51, 83, 19,  99, 35, 67, 3,
45    125, 61, 93, 29, 109, 45, 77, 13, 117, 53, 85, 21, 101, 37, 69, 5,
46    121, 57, 89, 25, 105, 41, 73,  9, 113, 49, 81, 17,  97, 33, 65, 1,
47];
48
49#[derive(Debug, Clone, Default, Encode, Decode)]
50pub struct Sdd1Decompressor {
51    source_addr: u32,
52    input: u16,
53    plane: u8,
54    num_planes: u8,
55    y_location: u8,
56    valid_bits: i8,
57    high_context_bits: u16,
58    low_context_bits: u16,
59    bit_counter: [u16; 8],
60    prev_bits: [u16; 8],
61    context_states: [u8; 32],
62    context_mps: [u8; 32],
63}
64
65impl Sdd1Decompressor {
66    pub fn new() -> Self {
67        Self::default()
68    }
69
70    pub fn init(&mut self, source_address: u32, mmc: &Sdd1Mmc, rom: &[u8]) {
71        self.input = read_byte(source_address, mmc, rom).into();
72        self.source_addr = source_address + 1;
73
74        self.num_planes = match self.input & 0xC0 {
75            // 2bpp tile data
76            0x00 => 2,
77            // 8bpp tile data
78            0x40 => 8,
79            // 4bpp tile data
80            0x80 => 4,
81            // Other data (e.g. Mode 7 graphics)
82            0xC0 => 0,
83            _ => unreachable!("value & 0xC0 is always one of the above values"),
84        };
85
86        // Context is formed using 3 or 4 of the previous 9 bits, with separate contexts for
87        // even and odd bitplanes
88        let (high_context_bits, low_context_bits) = match self.input & 0x30 {
89            // Bits 1, 7, 8, 9
90            0x00 => (0x01C0, 0x0001),
91            // Bits 1, 8, 9
92            0x10 => (0x0180, 0x0001),
93            // Bits 1, 7, 8
94            0x20 => (0x00C0, 0x0001),
95            // Bits 1, 2, 8, 9
96            0x30 => (0x0180, 0x0003),
97            _ => unreachable!("value & 0x30 is always one of the above values"),
98        };
99        self.high_context_bits = high_context_bits;
100        self.low_context_bits = low_context_bits;
101
102        let next_byte: u16 = read_byte(self.source_addr, mmc, rom).into();
103        self.input = (self.input << 11) | (next_byte << 3);
104        self.source_addr += 1;
105
106        self.valid_bits = 5;
107
108        self.bit_counter.fill(0);
109        self.prev_bits.fill(0);
110        self.context_states.fill(0);
111        self.context_mps.fill(0);
112
113        self.plane = 0;
114        self.y_location = 0;
115    }
116
117    pub fn next_byte(&mut self, mmc: &Sdd1Mmc, rom: &[u8]) -> u8 {
118        if self.num_planes == 0 {
119            // For miscellaneous data, simply output the next 8 bits
120            let mut byte = 0;
121            for plane in 0..8 {
122                byte |= self.get_bit(plane, mmc, rom) << plane;
123            }
124            return byte;
125        }
126
127        if !self.plane.bit(0) {
128            // Retrieve the next 16 bits, alternating between the even bitplane and the odd bitplane
129            for _ in 0..8 {
130                self.get_bit(self.plane, mmc, rom);
131                self.get_bit(self.plane + 1, mmc, rom);
132            }
133
134            let byte = self.prev_bits[self.plane as usize] & 0xFF;
135            self.plane += 1;
136
137            byte as u8
138        } else {
139            let byte = self.prev_bits[self.plane as usize] & 0xFF;
140            self.plane -= 1;
141
142            self.y_location += 1;
143            if self.y_location == 8 {
144                // Completed a set of 16 bytes; move to the next 2 bitplanes (if 4bpp or 8bpp)
145                self.y_location = 0;
146                self.plane = (self.plane + 2) & (self.num_planes - 1);
147            }
148
149            byte as u8
150        }
151    }
152
153    fn get_bit(&mut self, plane: u8, mmc: &Sdd1Mmc, rom: &[u8]) -> u8 {
154        // Form context from previous bits in the current plane, with separate contexts for odd
155        // and even bitplanes
156        let mut context = (u16::from(plane) & 0x01) << 4;
157        context |= (self.prev_bits[plane as usize] & self.high_context_bits) >> 5;
158        context |= self.prev_bits[plane as usize] & self.low_context_bits;
159
160        let p_bit = self.get_probable_bit(context, mmc, rom);
161        self.prev_bits[plane as usize] = (self.prev_bits[plane as usize] << 1) | u16::from(p_bit);
162
163        p_bit
164    }
165
166    fn get_probable_bit(&mut self, context: u16, mmc: &Sdd1Mmc, rom: &[u8]) -> u8 {
167        let state = self.context_states[context as usize];
168        let code_size = EVOLUTION_CODE_SIZE[state as usize];
169
170        if self.bit_counter[code_size as usize] & 0x7F == 0 {
171            self.bit_counter[code_size as usize] = self.get_codeword(code_size, mmc, rom);
172        }
173
174        let mut p_bit = self.context_mps[context as usize];
175        self.bit_counter[code_size as usize] -= 1;
176
177        if self.bit_counter[code_size as usize] == 0x00 {
178            // Run ends in the LPS
179            self.context_states[context as usize] = EVOLUTION_LPS_NEXT[state as usize];
180            p_bit ^= 0x01;
181
182            if state < 2 {
183                // MPS can only change while in state 0 or 1
184                self.context_mps[context as usize] = p_bit;
185            }
186        } else if self.bit_counter[code_size as usize] == 0x80 {
187            // Run ends in the MPS
188            self.context_states[context as usize] = EVOLUTION_MPS_NEXT[state as usize];
189        }
190
191        p_bit
192    }
193
194    fn get_codeword(&mut self, code_size: u8, mmc: &Sdd1Mmc, rom: &[u8]) -> u16 {
195        if self.valid_bits == 0 {
196            // Read next input byte
197            self.input |= u16::from(read_byte(self.source_addr, mmc, rom));
198            self.source_addr += 1;
199            self.valid_bits = 8;
200        }
201
202        self.input <<= 1;
203        self.valid_bits -= 1;
204
205        if !self.input.bit(15) {
206            // 0 indicates a run of MPSs of length 2^N, where N is the codeword size
207            return 0x80 + (1 << code_size);
208        }
209
210        // 1 indicates a run of MPSs that ends with the LPS, where the following N bits determine
211        // the run length
212
213        let run_table_idx = ((self.input >> 8) & 0x7F) | (0x7F >> code_size);
214        self.input <<= code_size;
215        self.valid_bits -= code_size as i8;
216        if self.valid_bits < 0 {
217            let next_byte: u16 = read_byte(self.source_addr, mmc, rom).into();
218            self.input |= next_byte << (-self.valid_bits);
219            self.source_addr += 1;
220            self.valid_bits += 8;
221        }
222
223        RUN_TABLE[run_table_idx as usize].into()
224    }
225}
226
227fn read_byte(address: u32, mmc: &Sdd1Mmc, rom: &[u8]) -> u8 {
228    mmc
229        .map_rom_address(address, rom.len() as u32)
230        .and_then(|rom_addr| rom.get(rom_addr as usize).copied())
231        .unwrap_or_else(|| {
232            log::error!("Encountered an invalid ROM address mapping in S-DD1 decompressor ({address:06X}); something has likely gone horribly wrong");
233            0
234        })
235}