1use crate::num::GetBit;
2use bincode::{Decode, Encode};
3use rustc_hash::FxHashMap;

End address inclusive

6#[derive(Debug, Clone, Encode, Decode)]
7pub struct CheatWordOverrides<const START_ADDRESS: usize, const END_ADDRESS: usize> {
8    address_bitset: Box<[u64]>,
9    memory_overrides: FxHashMap<u32, u16>,
10}
12impl<const START_ADDRESS: usize, const END_ADDRESS: usize>
13    CheatWordOverrides<START_ADDRESS, END_ADDRESS>
14{
15    #[must_use]
16    pub fn new(cheat_codes: &[(u32, u16)]) -> Self {
17        let mut overrides = Self {
18            address_bitset: vec![0; Self::bitset_len()].into_boxed_slice(),
19            memory_overrides: FxHashMap::default(),
20        };
21
22        overrides.update_cheat_codes(cheat_codes);
23        overrides
24    }
25
26    pub fn update_cheat_codes(&mut self, cheat_codes: &[(u32, u16)]) {
27        self.address_bitset.fill(0);
28        self.memory_overrides.clear();
29
30        for &(address, value) in cheat_codes {
31            if !(START_ADDRESS..=END_ADDRESS).contains(&(address as usize)) {
32                continue;
33            }
34
35            self.address_bitset[Self::address_to_bitset_idx(address)] |=
36                1 << Self::address_to_bitset_bit(address);
37            self.memory_overrides.insert(address & !1, value);
38        }
39
40        if !self.memory_overrides.is_empty() {
41            log::debug!("Cheat codes: {:X?}", self.memory_overrides);
42        }
43    }
44
45    #[must_use]
46    pub fn get(&self, address: u32) -> Option<u16> {
47        if self.memory_overrides.is_empty() {
48            return None;
49        }
50
51        if !self
52            .address_bitset
53            .get(Self::address_to_bitset_idx(address))
54            .is_some_and(|&bits| bits.bit(Self::address_to_bitset_bit(address)))
55        {
56            return None;
57        }
58
59        self.memory_overrides.get(&(address & !1)).copied()
60    }
61
62    const fn address_range_words() -> usize {
63        ((END_ADDRESS - START_ADDRESS) >> 1) + 1
64    }
65
66    const fn bitset_len() -> usize {
67        Self::address_range_words().div_ceil(64)
68    }
69
70    const fn address_to_bitset_idx(address: u32) -> usize {
71        let address = address as usize;
72        ((address.wrapping_sub(START_ADDRESS)) >> 1) / 64
73    }
74
75    const fn address_to_bitset_bit(address: u32) -> u8 {
76        ((address >> 1) & 63) as u8
77    }
78}
79
80#[derive(Debug, Clone, Copy, PartialEq, Eq, Encode, Decode)]
81pub struct ByteCheatCodeU16Address {
82    pub address: u16,
83    pub value: u8,
84    // If present, only override reads when the value in memory matches the reference value
85    pub reference: Option<u8>,
86}
87
88#[derive(Debug, Clone, Encode, Decode)]
89pub struct CheatByteOverridesU16Address {
90    address_bitset: Box<[u64]>,
91    overrides: FxHashMap<u16, (u8, Option<u8>)>,
92}
93
94impl CheatByteOverridesU16Address {
95    #[must_use]
96    pub fn new(cheat_codes: &[ByteCheatCodeU16Address]) -> Self {
97        let mut overrides = Self {
98            address_bitset: vec![0; 0x10000 / 64].into_boxed_slice(),
99            overrides: FxHashMap::default(),
100        };
101
102        overrides.update_cheat_codes(cheat_codes);
103        overrides
104    }
105
106    pub fn update_cheat_codes(&mut self, cheat_codes: &[ByteCheatCodeU16Address]) {
107        self.address_bitset.fill(0);
108        self.overrides.clear();
109
110        for &ByteCheatCodeU16Address { address, value, reference } in cheat_codes {
111            self.address_bitset[(address / 64) as usize] |= 1 << (address & 63);
112            self.overrides.insert(address, (value, reference));
113        }
114    }
115
116    #[must_use]
117    pub fn get(&self, address: u16, memory_value: u8) -> Option<u8> {
118        if self.overrides.is_empty() {
119            return None;
120        }
121
122        if !self.address_bitset.get((address / 64) as usize)?.bit((address & 63) as u8) {
123            return None;
124        }
125
126        let (override_value, reference) = *self.overrides.get(&address)?;
127        if reference.is_some_and(|reference| reference != memory_value) {
128            return None;
129        }
130
131        Some(override_value)
132    }
133}
134
135#[cfg(test)]
136mod tests {
137    use super::*;
138
139    #[test]
140    fn word_overrides() {
141        type TestOverrides = CheatWordOverrides<0xFF0000, 0xFFFFFF>;
142
143        let mut overrides = TestOverrides::new(&[]);
144
145        overrides.update_cheat_codes(&[(0xFFFFFF, 0x1234)]);
146        assert_eq!(overrides.get(0xFFFFFE), Some(0x1234));
147        assert_eq!(overrides.get(0xFFFFFF), Some(0x1234));
148        for address in 0xFF0000..=0xFF00FF {
149            assert_eq!(overrides.get(address), None);
150        }
151        assert_eq!(overrides.get(0x0000FF), None);
152
153        overrides.update_cheat_codes(&[(0x001234, 0x5678), (0x1000000, 0xABCD)]);
154        assert_eq!(overrides.get(0x001234), None);
155        assert_eq!(overrides.get(0xFFFFFF), None);
156        assert_eq!(overrides.get(0x1000000), None);
157    }
158}