1use crate::ppu::registers::{BitsPerPixel, ObjPriorityMode, TileSize};
2use crate::ppu::{
3    MAX_SPRITE_TILES_PER_LINE, MAX_SPRITES_PER_LINE, Pixel, Ppu, VRAM_ADDRESS_MASK,
4    line_overlaps_sprite,
5};
6use bincode::{Decode, Encode};
7use jgenesis_common::num::GetBit;
8use std::cmp;
9
10pub const SPRITE_EVALUATION_END_DOT: u16 = 256;

340 - 270 = 70 Tile limit is 34 (68/2), but give 2 extra dots so that it's possible for a 35th tile to trigger the sprite time overflow check

15pub const SPRITE_FETCH_START_DOT: u16 = 270;
16pub const SPRITE_FETCH_END_DOT: u16 = 340;
18#[derive(Debug, Clone, Copy, PartialEq, Eq, Encode, Decode)]
19pub enum SpriteState {
20    Blank,
21    Evaluation { oam_idx: u8 },
22    Idle { oam_idx: u8 },
23    TileFetch { oam_buffer_idx: u8, tile_idx: u8 },
24}
25
26impl SpriteState {
27    pub fn oam_idx(self, scanned_oam_idxs: &[u8]) -> u8 {
28        match self {
29            Self::Evaluation { oam_idx } | Self::Idle { oam_idx } => oam_idx,
30            Self::TileFetch { oam_buffer_idx, .. } => scanned_oam_idxs[oam_buffer_idx as usize],
31            Self::Blank => 0,
32        }
33    }
34}
35
36#[derive(Debug, Clone, Encode, Decode)]
37pub struct SpriteTileData {
38    pub x: u16,
39    pub palette: u8,
40    pub priority: u8,
41    pub colors: [u8; 8],
42}
43
44#[derive(Debug, Clone, Encode, Decode)]
45pub struct SpriteProcessor {
46    pub state: SpriteState,
47    pub line: u16,
48    pub interlaced_odd_frame: bool,
49    pub dot: u16,
50    pub scanned_oam_idxs: Vec<u8>,
51    pub fetched_tiles: Vec<SpriteTileData>,
52    pub fetched_tiles_deinterlace: Vec<SpriteTileData>,
53    pub last_fetched_oam_idx: u8,
54}
55
56impl SpriteProcessor {
57    pub fn new() -> Self {
58        Self {
59            state: SpriteState::Blank,
60            line: 0,
61            interlaced_odd_frame: false,
62            dot: 0,
63            scanned_oam_idxs: Vec::with_capacity(MAX_SPRITES_PER_LINE),
64            fetched_tiles: Vec::with_capacity(MAX_SPRITE_TILES_PER_LINE),
65            fetched_tiles_deinterlace: Vec::with_capacity(MAX_SPRITE_TILES_PER_LINE),
66            last_fetched_oam_idx: 0,
67        }
68    }
69}
70
71impl Ppu {
72    pub(super) fn sprites_start_new_line(&mut self, scanline: u16, interlaced_odd_frame: bool) {
73        self.sprites.line = scanline;
74        self.sprites.interlaced_odd_frame = interlaced_odd_frame;
75        self.sprites.dot = 0;
76
77        self.sprites.state = if self.in_active_display(scanline) {
78            // If priority rotate mode is set, start iteration at the current OAM address instead of 0
79            self.sprites.scanned_oam_idxs.clear();
80            let start_oam_idx = match self.registers.obj_priority_mode {
81                ObjPriorityMode::Normal => 0,
82                ObjPriorityMode::Rotate => (self.registers.oam_address >> 1) & 0x7F,
83            };
84
85            log::trace!(
86                "Beginning sprite evaluation for line {scanline} at OAM idx {start_oam_idx}"
87            );
88
89            SpriteState::Evaluation { oam_idx: start_oam_idx as u8 }
90        } else {
91            // Sprite evaluation not performed in VBlank or forced blanking
92            SpriteState::Blank
93        };
94    }
95
96    pub(super) fn progress_sprite_evaluation(&mut self, dot: u16) {
97        let SpriteState::Evaluation { oam_idx } = self.sprites.state else { return };
98
99        let dot = cmp::min(dot, SPRITE_EVALUATION_END_DOT);
100
101        // Sprite evaluation takes 2 dots per OAM entry
102        debug_assert!(self.sprites.dot <= dot);
103        let num_sprites_to_scan = dot / 2 - self.sprites.dot / 2;
104
105        log::trace!(
106            "Progressing sprite evaluation on line {} to dot {dot}, scanning up to {num_sprites_to_scan} sprites",
107            self.sprites.line
108        );
109
110        let new_oam_idx = self.progress_oam_scan(self.sprites.line, oam_idx, num_sprites_to_scan);
111
112        self.sprites.dot = dot;
113        self.sprites.state = if dot == SPRITE_EVALUATION_END_DOT
114            || self.sprites.scanned_oam_idxs.len() == MAX_SPRITES_PER_LINE
115        {
116            SpriteState::Idle {
117                oam_idx: if self.sprites.scanned_oam_idxs.is_empty() {
118                    // If no sprites were scanned in range for this line, mid-scanline OAM writes
119                    // should go to the sprite for the last fetched tile.
120                    // Uniracers depends on this for correct rendering in Vs. mode
121                    self.sprites.last_fetched_oam_idx
122                } else {
123                    0
124                },
125            }
126        } else {
127            SpriteState::Evaluation { oam_idx: new_oam_idx }
128        };
129    }
130
131    #[must_use]
132    fn progress_oam_scan(
133        &mut self,
134        scanline: u16,
135        start_oam_idx: u8,
136        num_sprites_to_scan: u16,
137    ) -> u8 {
138        const OAM_IDX_MASK: u8 = 0x7F;
139
140        let (small_width, small_height, large_width, large_height) = {
141            let (small_width, mut small_height) = self.registers.obj_tile_size.small_size();
142            let (large_width, mut large_height) = self.registers.obj_tile_size.large_size();
143
144            if self.registers.interlaced && self.registers.pseudo_obj_hi_res {
145                // If smaller OBJs are enabled, pretend sprites are half-size vertically for the OAM scan
146                small_height >>= 1;
147                large_height >>= 1;
148            }
149
150            (small_width, small_height, large_width, large_height)
151        };
152
153        let mut oam_idx = start_oam_idx;
154        for _ in 0..num_sprites_to_scan {
155            let oam_low_addr = (oam_idx << 1) as usize;
156            let [x_lsb, y] = self.oam_low[oam_low_addr].to_le_bytes();
157
158            let oam_high_addr = (oam_idx >> 2) as usize;
159            let oam_high_shift = 2 * (oam_idx & 3);
160            let oam_high_bits = self.oam_high[oam_high_addr] >> oam_high_shift;
161
162            let x_msb = oam_high_bits.bit(0);
163            let size = if oam_high_bits.bit(1) { TileSize::Large } else { TileSize::Small };
164
165            let (sprite_width, sprite_height) = match size {
166                TileSize::Small => (small_width, small_height),
167                TileSize::Large => (large_width, large_height),
168            };
169
170            if !line_overlaps_sprite(y, sprite_height, scanline) {
171                oam_idx = (oam_idx + 1) & OAM_IDX_MASK;
172                continue;
173            }
174
175            // Only sprites with pixels in the range [0, 256) are scanned into the sprite buffer
176            let x = u16::from_le_bytes([x_lsb, u8::from(x_msb)]);
177            if x >= 256 && x + sprite_width <= 512 {
178                oam_idx = (oam_idx + 1) & OAM_IDX_MASK;
179                continue;
180            }
181
182            if self.sprites.scanned_oam_idxs.len() == MAX_SPRITES_PER_LINE {
183                self.registers.sprite_overflow = true;
184                log::debug!("Hit 32 sprites per line limit on line {scanline}");
185                return (oam_idx + 1) & OAM_IDX_MASK;
186            }
187
188            self.sprites.scanned_oam_idxs.push(oam_idx);
189            oam_idx = (oam_idx + 1) & OAM_IDX_MASK;
190        }
191
192        oam_idx
193    }
194
195    pub(super) fn begin_sprite_tile_fetch(&mut self) {
196        if self.vblank_flag() {
197            // No sprite tile fetching during VBlank
198            self.sprites.state = SpriteState::Blank;
199            return;
200        }
201
202        // Explicitly do not check whether forced blanking is enabled.
203        // If forced blanking is enabled, let the tile fetch proceed as normal, but make any fetched
204        // sprites invisible (all pixels transparent).
205
206        self.sprites.dot = SPRITE_FETCH_START_DOT;
207        self.sprites.fetched_tiles.clear();
208
209        if self.sprites.scanned_oam_idxs.is_empty() {
210            // No tiles to fetch
211            log::trace!("No sprites scanned during evaluation for line {}", self.sprites.line);
212            self.sprites.state = SpriteState::Idle { oam_idx: self.sprites.last_fetched_oam_idx };
213            return;
214        }
215
216        log::trace!(
217            "Beginning sprite tile fetch for line {}, scanned {} sprites",
218            self.sprites.line,
219            self.sprites.scanned_oam_idxs.len()
220        );
221
222        // Tiles are fetched for sprites in reverse order (games depend on this, e.g. Final Fantasy 6)
223        // Tiles within a sprite are processed left-to-right
224        self.sprites.state = SpriteState::TileFetch {
225            oam_buffer_idx: (self.sprites.scanned_oam_idxs.len() - 1) as u8,
226            tile_idx: 0,
227        };
228    }
229
230    pub(super) fn progress_sprite_tile_fetch(&mut self, dot: u16) {
231        if !matches!(self.sprites.state, SpriteState::TileFetch { .. }) {
232            return;
233        }
234
235        let end_dot = cmp::min(dot, SPRITE_FETCH_END_DOT);
236        if self.sprites.dot >= end_dot {
237            return;
238        }
239
240        // Tiles are fetched at a rate of 2 dots per tile
241        let num_tiles_to_fetch = end_dot / 2 - self.sprites.dot / 2;
242
243        log::trace!(
244            "Progressing sprite tile fetch for line {} to dot {dot}, fetching up to {num_tiles_to_fetch} tiles",
245            self.sprites.line
246        );
247
248        self.sprites.state = self.fetch_sprite_tiles(
249            self.sprites.line,
250            end_dot,
251            self.sprites.interlaced_odd_frame,
252            num_tiles_to_fetch,
253        );
254        self.sprites.dot = end_dot;
255    }
256
257    pub(super) fn sprites_finish_line(&mut self) {
258        self.progress_sprite_tile_fetch(SPRITE_FETCH_END_DOT);
259
260        log::trace!(
261            "Rendering {} sprite tiles to line buffer for line {}",
262            self.sprites.fetched_tiles.len(),
263            self.sprites.line
264        );
265
266        self.render_sprite_tiles();
267
268        if self.deinterlace
269            && self.state.v_hi_res_frame
270            && self.registers.pseudo_obj_hi_res
271            && !self.sprites.scanned_oam_idxs.is_empty()
272        {
273            log::trace!("Fetching extra line of sprite tiles for deinterlaced rendering");
274
275            self.sprites.fetched_tiles.clear();
276            self.sprites.state = SpriteState::TileFetch {
277                oam_buffer_idx: (self.sprites.scanned_oam_idxs.len() - 1) as u8,
278                tile_idx: 0,
279            };
280            self.sprites.state = self.fetch_sprite_tiles(
281                self.sprites.line,
282                SPRITE_FETCH_END_DOT,
283                !self.sprites.interlaced_odd_frame,
284                MAX_SPRITE_TILES_PER_LINE as u16,
285            );
286        }
287    }
288
289    #[must_use]
290    fn fetch_sprite_tiles(
291        &mut self,
292        scanline: u16,
293        dot: u16,
294        interlaced_odd_line: bool,
295        num_tiles_to_fetch: u16,
296    ) -> SpriteState {
297        let SpriteState::TileFetch { mut oam_buffer_idx, mut tile_idx } = self.sprites.state else {
298            return self.sprites.state;
299        };
300
301        let (small_width, small_height) = self.registers.obj_tile_size.small_size();
302        let (large_width, large_height) = self.registers.obj_tile_size.large_size();
303
304        let mut tiles_fetched = 0;
305        while tiles_fetched < num_tiles_to_fetch {
306            let oam_idx = self.sprites.scanned_oam_idxs[oam_buffer_idx as usize];
307
308            let oam_low_addr = usize::from(oam_idx) << 1;
309            let [x_lsb, y] = self.oam_low[oam_low_addr].to_le_bytes();
310
311            let [tile_number_lsb, attributes] = self.oam_low[oam_low_addr + 1].to_le_bytes();
312
313            let oam_high_addr = usize::from(oam_idx >> 2);
314            let oam_high_shift = 2 * (oam_idx & 3);
315            let oam_high_bits = self.oam_high[oam_high_addr] >> oam_high_shift;
316
317            let base_tile_number =
318                u16::from_le_bytes([tile_number_lsb, u8::from(attributes.bit(0))]);
319            let palette = (attributes >> 1) & 0x07;
320            let priority = (attributes >> 4) & 0x03;
321            let x_flip = attributes.bit(6);
322            let y_flip = attributes.bit(7);
323
324            let x_msb = oam_high_bits.bit(0);
325            let size = if oam_high_bits.bit(1) { TileSize::Large } else { TileSize::Small };
326
327            let x = u16::from_le_bytes([x_lsb, x_msb.into()]);
328
329            let (sprite_width, sprite_height) = match size {
330                TileSize::Small => (small_width, small_height),
331                TileSize::Large => (large_width, large_height),
332            };
333
334            if !line_overlaps_sprite(y, sprite_height, scanline) {
335                // Can happen if Y coordinate changes between OAM scan and tile fetch
336                if oam_buffer_idx == 0 {
337                    // Fetched all tiles for this line
338                    return SpriteState::Idle { oam_idx: self.sprites.last_fetched_oam_idx };
339                }
340
341                oam_buffer_idx -= 1;
342                tile_idx = 0;
343                continue;
344            }
345
346            let mut sprite_line = if y_flip {
347                sprite_height as u8
348                    - 1
349                    - ((scanline as u8).wrapping_sub(y) & ((sprite_height - 1) as u8))
350            } else {
351                (scanline as u8).wrapping_sub(y) & ((sprite_height - 1) as u8)
352            };
353
354            // Adjust sprite line if smaller OBJs are enabled
355            // Smaller OBJs affect how the line within the sprite is determined, but not where the
356            // sprite is positioned onscreen
357            if self.registers.interlaced && self.registers.pseudo_obj_hi_res {
358                sprite_line = (sprite_line << 1) | u8::from(interlaced_odd_line ^ y_flip);
359            }
360
361            let tile_y_offset: u16 = (sprite_line / 8).into();
362
363            let num_sprite_tiles = (sprite_width / 8) as u8;
364            while tile_idx < num_sprite_tiles && tiles_fetched < num_tiles_to_fetch {
365                let tile_x_offset: u16 = tile_idx.into();
366                let x = if x_flip {
367                    x + (sprite_width - 8) - 8 * tile_x_offset
368                } else {
369                    x + 8 * tile_x_offset
370                };
371
372                if x >= 256 && x + 8 < 512 {
373                    // Sprite tile is entirely offscreen; don't fetch
374                    tile_idx += 1;
375                    continue;
376                }
377
378                if self.sprites.fetched_tiles.len() == MAX_SPRITE_TILES_PER_LINE {
379                    // Sprite time overflow
380                    self.registers.sprite_pixel_overflow = true;
381                    log::debug!("Hit 34 sprite tiles per line limit on line {scanline}");
382                    return SpriteState::Idle { oam_idx: self.sprites.last_fetched_oam_idx };
383                }
384
385                // Unlike BG tiles in 16x16 mode, overflows in large OBJ tiles do not carry to the next nibble
386                let mut tile_number = base_tile_number;
387                tile_number =
388                    (tile_number & !0xF) | (tile_number.wrapping_add(tile_x_offset) & 0xF);
389                tile_number =
390                    (tile_number & !0xF0) | (tile_number.wrapping_add(tile_y_offset << 4) & 0xF0);
391
392                let tile_size_words = BitsPerPixel::OBJ.tile_size_words();
393                let tile_base_addr = self.registers.obj_tile_base_address
394                    + u16::from(tile_number.bit(8))
395                        * (256 * tile_size_words + self.registers.obj_tile_gap_size);
396                let tile_addr = ((tile_base_addr + (tile_number & 0x00FF) * tile_size_words)
397                    & VRAM_ADDRESS_MASK) as usize;
398
399                let tile_data = &self.vram[tile_addr..tile_addr + tile_size_words as usize];
400
401                let tile_row: u16 = (sprite_line % 8).into();
402
403                let mut colors = [0_u8; 8];
404                for tile_col in 0..8 {
405                    let bit_index = (7 - tile_col) as u8;
406
407                    let mut color = 0_u8;
408                    for i in 0..2 {
409                        let tile_word = tile_data[(tile_row + 8 * i) as usize];
410                        color |= u8::from(tile_word.bit(bit_index)) << (2 * i);
411                        color |= u8::from(tile_word.bit(bit_index + 8)) << (2 * i + 1);
412                    }
413
414                    colors[if x_flip { 7 - tile_col } else { tile_col }] = color;
415                }
416
417                self.sprites.fetched_tiles.push(if !self.registers.forced_blanking {
418                    SpriteTileData { x, palette, priority, colors }
419                } else {
420                    SpriteTileData { x: 0, palette: 0, priority: 0, colors: [0; 8] }
421                });
422                self.sprites.last_fetched_oam_idx = oam_idx;
423
424                tile_idx += 1;
425                tiles_fetched += 1;
426            }
427
428            if tile_idx == num_sprite_tiles {
429                if oam_buffer_idx == 0 {
430                    // Fetched all sprite tiles
431                    return SpriteState::Idle { oam_idx: self.sprites.last_fetched_oam_idx };
432                }
433
434                oam_buffer_idx -= 1;
435                tile_idx = 0;
436            }
437        }
438
439        if dot >= SPRITE_FETCH_END_DOT {
440            SpriteState::Idle { oam_idx: self.sprites.last_fetched_oam_idx }
441        } else {
442            SpriteState::TileFetch { oam_buffer_idx, tile_idx }
443        }
444    }
445
446    pub(super) fn render_sprite_tiles(&mut self) {
447        if self.vblank_flag()
448            || !(self.registers.main_obj_enabled || self.registers.sub_obj_enabled)
449        {
450            return;
451        }
452
453        self.buffers.obj_pixels.fill(Pixel::TRANSPARENT);
454        for tile in &self.sprites.fetched_tiles {
455            for dx in 0..8 {
456                let x = (tile.x + dx) & 0x1FF;
457                if x >= 256 {
458                    continue;
459                }
460
461                let pixel_color = tile.colors[dx as usize];
462                if pixel_color == 0 {
463                    // Transparent
464                    continue;
465                }
466
467                self.buffers.obj_pixels[x as usize] =
468                    Pixel { palette: tile.palette, color: pixel_color, priority: tile.priority };
469            }
470        }
471    }
472
473    pub(super) fn progress_for_mid_scanline_write(&mut self, dot: u16) {
474        log::debug!(
475            "Progressing sprite state to dot {dot} on line {} for active display OAMDATA/INIDISP write",
476            self.sprites.line
477        );
478
479        match self.sprites.state {
480            SpriteState::Evaluation { .. } => {
481                self.progress_sprite_evaluation(dot);
482            }
483            SpriteState::TileFetch { .. } => {
484                self.progress_sprite_tile_fetch(dot);
485            }
486            _ => {}
487        }
488    }
489
490    pub(super) fn sprites_forced_blanking_change(&mut self, new_forced_blanking: bool) {
491        if self.vblank_flag() {
492            // Changing forced blanking during VBlank has no effect on sprite state
493            return;
494        }
495
496        log::debug!(
497            "Handling sprite state update for mid-scanline forced blanking change to {new_forced_blanking} on line {}",
498            self.sprites.line
499        );
500
501        let dot = (self.state.scanline_master_cycles / 4) as u16;
502        self.progress_for_mid_scanline_write(dot);
503        self.sprites.dot = dot;
504
505        if new_forced_blanking {
506            // Reset OAM address based on current sprite evaluation/fetching state
507            self.registers.oam_address =
508                (self.sprites.state.oam_idx(&self.sprites.scanned_oam_idxs) << 1).into();
509            if !matches!(self.sprites.state, SpriteState::TileFetch { .. }) {
510                self.sprites.state = SpriteState::Blank;
511            }
512            return;
513        }
514
515        let oam_idx = ((self.registers.oam_address >> 1) & 0x7F) as u8;
516        match dot {
517            0..=255 => {
518                // Disabling forced blanking at H<256 causes the PPU to resume sprite evaluation from
519                // the current OAM address
520                self.sprites.state = SpriteState::Evaluation { oam_idx };
521            }
522            270..=u16::MAX => {
523                // Fetching tiles; don't update state
524            }
525            _ => {
526                self.sprites.state = SpriteState::Idle { oam_idx };
527            }
528        }
529    }
530}