1use crate::mainloop::bincode_config;
2use bincode::{Decode, Encode};
3use jgenesis_common::frontend::{EmulatorTrait, Renderer};
4use std::collections::VecDeque;
5use std::sync::atomic::{AtomicI32, Ordering};
6use std::sync::mpsc::{Receiver, Sender, SyncSender};
7use std::sync::{Arc, mpsc};
8use std::thread;
9use std::time::{Duration, Instant};
10
11const FRAME_DIVIDER: u64 = 10;
12const REWIND_SPEED: u64 = 2;
13
14struct CompressorThread<State> {
15    state_receiver: Receiver<State>,
16    bytes_sender: Sender<Vec<u8>>,
17    requests_in_flight: Arc<AtomicI32>,
18}
19
20struct CompressorThreadHandle<Emulator: EmulatorTrait> {
21    state_sender: SyncSender<Box<Emulator::SaveState>>,
22    bytes_receiver: Receiver<Vec<u8>>,
23    requests_in_flight: Arc<AtomicI32>,
24}
25
26impl<Emulator: EmulatorTrait> CompressorThreadHandle<Emulator> {
27    fn send_state(&self, state: Box<Emulator::SaveState>) {
28        if self.state_sender.send(state).is_err() {
29            log::error!("Lost connection to rewind compression thread; this is probably a bug");
30            return;
31        }
32
33        self.requests_in_flight.fetch_add(1, Ordering::AcqRel);
34    }
35
36    fn try_recv_compressed_bytes(&self) -> Option<Vec<u8>> {
37        self.bytes_receiver.try_recv().ok()
38    }
39
40    fn any_requests_in_flight(&self) -> bool {
41        self.requests_in_flight.load(Ordering::Acquire) > 0
42    }
43}
44
45fn spawn_compressor_thread<Emulator: EmulatorTrait>() -> CompressorThreadHandle<Emulator> {
46    // Bound state sender channel to prevent it from growing infinitely if compressor can't keep up
47    let (state_sender, state_receiver) = mpsc::sync_channel(10);
48    let (bytes_sender, bytes_receiver) = mpsc::channel();
49    let requests_in_flight = Arc::new(AtomicI32::new(0));
50
51    let compressor_thread = CompressorThread {
52        state_receiver,
53        bytes_sender,
54        requests_in_flight: Arc::clone(&requests_in_flight),
55    };
56
57    thread::spawn(move || run_compressor_thread(compressor_thread));
58
59    CompressorThreadHandle { state_sender, bytes_receiver, requests_in_flight }
60}
61
62fn run_compressor_thread<State: Encode>(compressor: CompressorThread<State>) {
63    loop {
64        let Ok(state) = compressor.state_receiver.recv() else {
65            // Runner thread has dropped sender; stop running
66            return;
67        };
68
69        let compressed_bytes = compress_state(&state);
70        if let Some(compressed_bytes) = compressed_bytes
71            && compressor.bytes_sender.send(compressed_bytes).is_err()
72        {
73            // Runner thread has dropped receiver; stop running
74            return;
75        }
76
77        compressor.requests_in_flight.fetch_sub(1, Ordering::AcqRel);
78    }
79}
80
81fn compress_state<State: Encode>(state: &State) -> Option<Vec<u8>> {
82    let mut encoder = zstd::Encoder::new(vec![], 0).ok()?;
83    bincode::encode_into_std_write(state, &mut encoder, bincode_config!()).ok()?;
84    encoder.finish().ok()
85}
86
87fn decompress_state<State: Decode<()>>(bytes: &[u8]) -> Option<State> {
88    let mut decoder = zstd::Decoder::new(bytes).ok()?;
89    bincode::decode_from_std_read(&mut decoder, bincode_config!()).ok()
90}
91
92pub struct Rewinder<Emulator: EmulatorTrait> {
93    previous_states: VecDeque<Vec<u8>>,
94    buffer_len: usize,
95    frame_count: u64,
96    interval_multiplier: f64,
97    last_rewind_time: Option<Instant>,
98    compressor_handle: CompressorThreadHandle<Emulator>,
99}
100
101impl<Emulator: EmulatorTrait> Rewinder<Emulator> {
102    pub fn new(buffer_duration: Duration) -> Self {
103        let buffer_len = duration_to_buffer_len(buffer_duration);
104        let compressor_handle = spawn_compressor_thread();
105
106        Self {
107            previous_states: VecDeque::with_capacity(buffer_len + 1),
108            buffer_len,
109            frame_count: 0,
110            interval_multiplier: 1.0,
111            last_rewind_time: None,
112            compressor_handle,
113        }
114    }
115
116    pub fn record_frame(&mut self, emulator: &Emulator) {
117        if self.buffer_len == 0 {
118            return;
119        }
120
121        self.frame_count += 1;
122
123        if self.frame_count.is_multiple_of(FRAME_DIVIDER) {
124            self.compressor_handle.send_state(Box::new(emulator.to_save_state()));
125        }
126
127        self.recv_queued_compressed_bytes();
128    }
129
130    fn recv_queued_compressed_bytes(&mut self) {
131        while let Some(compressed_bytes) = self.compressor_handle.try_recv_compressed_bytes() {
132            self.previous_states.push_back(compressed_bytes);
133
134            while self.previous_states.len() > self.buffer_len {
135                self.previous_states.pop_front();
136            }
137        }
138    }
139
140    pub fn start_rewinding(&mut self) {
141        if self.last_rewind_time.is_none() {
142            self.last_rewind_time = Some(Instant::now());
143        }
144
145        while self.compressor_handle.any_requests_in_flight() {
146            thread::sleep(Duration::from_millis(1));
147        }
148
149        self.recv_queued_compressed_bytes();
150    }
151
152    pub fn stop_rewinding(&mut self) {
153        self.last_rewind_time = None;
154    }
155
156    pub fn is_rewinding(&self) -> bool {
157        self.last_rewind_time.is_some()
158    }
159
160    pub fn tick<R>(
161        &mut self,
162        emulator: &mut Emulator,
163        renderer: &mut R,
164        config: &Emulator::Config,
165    ) -> Result<(), R::Err>
166    where
167        Emulator: EmulatorTrait,
168        R: Renderer,
169    {
170        let Some(last_rewind_time) = self.last_rewind_time else { return Ok(()) };
171
172        let rewind_interval_secs =
173            self.interval_multiplier / 60.0 * (FRAME_DIVIDER as f64) / (REWIND_SPEED as f64);
174
175        let now = Instant::now();
176        if now.duration_since(last_rewind_time) >= Duration::from_secs_f64(rewind_interval_secs) {
177            let Some(compressed_bytes) = self.previous_states.pop_back() else { return Ok(()) };
178
179            if let Some(state) = decompress_state(&compressed_bytes) {
180                emulator.load_state(state);
181
182                emulator.reload_config(config);
183                emulator.force_render(renderer)?;
184            }
185
186            self.last_rewind_time = Some(now);
187        }
188
189        Ok(())
190    }
191
192    pub fn set_buffer_duration(&mut self, duration: Duration) {
193        self.set_buffer_len(duration_to_buffer_len(duration));
194    }
195
196    fn set_buffer_len(&mut self, buffer_len: usize) {
197        self.buffer_len = buffer_len;
198
199        // If size increased, immediately resize deque to avoid incremental allocations later
200        if buffer_len + 1 > self.previous_states.capacity() {
201            self.previous_states.reserve(buffer_len + 1 - self.previous_states.capacity());
202        }
203
204        // If size decreased, immediately drop unused states
205        while self.previous_states.len() > buffer_len {
206            self.previous_states.pop_front();
207        }
208    }
209
210    pub fn set_speed_multiplier(&mut self, speed_multiplier: u64) {
211        self.interval_multiplier = match speed_multiplier {
212            1 => 1.0,
213            // Default rewind speed is 2x, so fudge 2x and 3x to make them between 2x and 4x speed
214            2 => 3.0 / 4.0,
215            3 => 3.0 / 5.0,
216            _ => (REWIND_SPEED as f64) / (speed_multiplier as f64),
217        };
218    }
219}
220
221fn duration_to_buffer_len(duration: Duration) -> usize {
222    // Not really a better place for this, and this should get optimized out anyway
223    assert_eq!(FRAME_DIVIDER % REWIND_SPEED, 0);
224
225    (duration.as_secs() * 60 / (FRAME_DIVIDER / REWIND_SPEED)) as usize
226}