use std::io::{self};
use std::io::{Read, Seek};
use thiserror::Error;
use crate::common::{
binpack_error::BinpackError, compressed_training_file_reader::CompressedTrainingDataFileReader,
entry::PackedTrainingDataEntry, entry::TrainingDataEntry,
};
use super::move_score_list_reader::PackedMoveScoreListReader;
const SUGGESTED_CHUNK_SIZE: usize = 8192;
#[derive(Debug, Error)]
pub enum CompressedReaderError {
#[error("IO error: {0}")]
Io(#[from] io::Error),
#[error("Invalid data format: {0}")]
InvalidFormat(String),
#[error("End of file reached")]
EndOfFile,
#[error("Binpack error: {0}")]
BinpackError(#[from] BinpackError),
}
type Result<T> = std::result::Result<T, CompressedReaderError>;
pub fn read_chunk_into<T: Read + Seek>(file: &mut T, buffer: &mut Vec<u8>) -> Result<bool> {
let mut reader = CompressedTrainingDataFileReader::new(file)?;
if !reader.has_next_chunk() {
return Ok(false);
}
reader.read_next_chunk_into(buffer)?;
Ok(true)
}
pub fn parse_chunk(chunk: &[u8]) -> Vec<TrainingDataEntry> {
let mut reader = ChunkReader::default();
let mut entries = Vec::new();
while reader.has_next(chunk) {
entries.push(reader.next(chunk));
}
entries
}
#[derive(Debug)]
pub struct CompressedTrainingDataEntryReader<T: Read + Seek> {
chunk: Vec<u8>,
chunk_reader: ChunkReader,
input_file: Option<CompressedTrainingDataFileReader<T>>,
is_end: bool,
}
#[derive(Debug, Default)]
pub struct ChunkReader {
movelist_reader: Option<PackedMoveScoreListReader>,
offset: usize,
is_end: bool,
}
impl<T: Read + Seek> CompressedTrainingDataEntryReader<T> {
pub fn new(file: T) -> Result<Self> {
let chunk = Vec::with_capacity(SUGGESTED_CHUNK_SIZE);
let mut reader = Self {
chunk,
chunk_reader: ChunkReader::default(),
input_file: Some(CompressedTrainingDataFileReader::new(file)?),
is_end: false,
};
if !reader.load_next_chunk()? {
reader.is_end = true;
return Err(CompressedReaderError::EndOfFile);
}
Ok(reader)
}
pub fn into_inner(&mut self) -> io::Result<T> {
self.input_file.take().unwrap().into_inner()
}
pub fn read_bytes(&self) -> u64 {
self.input_file.as_ref().unwrap().read_bytes()
}
pub fn read_next_chunk_into(&mut self, buffer: &mut Vec<u8>) -> Result<bool> {
if !self.input_file.as_mut().unwrap().has_next_chunk() {
return Ok(false);
}
self.input_file
.as_mut()
.unwrap()
.read_next_chunk_into(buffer)?;
Ok(true)
}
pub fn parse_chunk(chunk: &[u8]) -> Vec<TrainingDataEntry> {
parse_chunk(chunk)
}
pub fn has_next(&self) -> bool {
!self.is_end
}
pub fn is_next_entry_continuation(&self) -> bool {
if let Some(ref reader) = self.chunk_reader.movelist_reader {
return reader.has_next();
}
false
}
#[allow(clippy::should_implement_trait)]
pub fn next(&mut self) -> TrainingDataEntry {
let entry = self.chunk_reader.next(&self.chunk);
if !self.chunk_reader.has_next(&self.chunk) {
self.fetch_next_chunk_if_needed();
}
entry
}
fn fetch_next_chunk_if_needed(&mut self) {
if self.chunk_reader.has_next(&self.chunk) {
return;
}
if self.load_next_chunk().unwrap() {
return;
}
self.is_end = true;
}
fn load_next_chunk(&mut self) -> Result<bool> {
if !self.input_file.as_mut().unwrap().has_next_chunk() {
return Ok(false);
}
self.input_file
.as_mut()
.unwrap()
.read_next_chunk_into(&mut self.chunk)?;
self.chunk_reader = ChunkReader::default();
Ok(true)
}
}
impl ChunkReader {
pub fn has_next(&self, chunk: &[u8]) -> bool {
if self
.movelist_reader
.as_ref()
.is_some_and(|reader| reader.has_next())
{
return true;
}
!self.is_end && self.offset + PackedTrainingDataEntry::byte_size() + 2 <= chunk.len()
}
#[allow(clippy::should_implement_trait)]
pub fn next(&mut self, chunk: &[u8]) -> TrainingDataEntry {
if let Some(ref mut reader) = self.movelist_reader {
let entry = reader.next_entry(&chunk[self.offset..]);
if !reader.has_next() {
self.offset += reader.num_read_bytes();
self.movelist_reader = None;
self.finish_if_at_end(chunk);
}
return entry;
}
let entry = self.read_entry(chunk);
let num_plies = self.read_plies(chunk);
if num_plies > 0 {
self.movelist_reader = Some(PackedMoveScoreListReader::new(entry, num_plies));
} else {
self.finish_if_at_end(chunk);
}
entry
}
fn read_entry(&mut self, chunk: &[u8]) -> TrainingDataEntry {
let size = PackedTrainingDataEntry::byte_size();
debug_assert!(self.offset + size <= chunk.len());
let packed = PackedTrainingDataEntry::from_slice(&chunk[self.offset..self.offset + size]);
self.offset += size;
packed.unpack_entry()
}
fn read_plies(&mut self, chunk: &[u8]) -> u16 {
let ply = ((chunk[self.offset] as u16) << 8) | (chunk[self.offset + 1] as u16);
self.offset += 2;
ply
}
fn finish_if_at_end(&mut self, chunk: &[u8]) {
if self.offset + PackedTrainingDataEntry::byte_size() + 2 > chunk.len() {
self.is_end = true;
}
}
}
impl CompressedTrainingDataEntryReader<io::Cursor<Vec<u8>>> {
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
Self::new(io::Cursor::new(bytes))
}
}
impl<'a> CompressedTrainingDataEntryReader<io::Cursor<&'a [u8]>> {
pub fn from_slice(bytes: &'a [u8]) -> Result<Self> {
Self::new(io::Cursor::new(bytes))
}
}
#[cfg(test)]
mod tests {
use std::{fs::OpenOptions, io::Cursor};
use crate::chess::{
coords::Square,
piece::Piece,
position::Position,
r#move::{Move, MoveType},
};
use super::*;
#[test]
fn test_reader_simple() {
let file = OpenOptions::new()
.read(true)
.write(true)
.create(false)
.append(false)
.open("./test/ep1.binpack")
.unwrap();
let mut reader = CompressedTrainingDataEntryReader::new(file).unwrap();
let mut entries: Vec<TrainingDataEntry> = Vec::new();
while reader.has_next() {
let entry = reader.next();
entries.push(entry);
}
let expected = vec![
TrainingDataEntry {
pos: Position::from_fen("1q5b/1r5k/4p2p/1b2P1pN/3p4/6PP/1nP3B1/1Q2B1K1 w - - 0 35")
.unwrap(),
mv: Move::new(
Square::new(10),
Square::new(26),
MoveType::Normal,
Piece::none(),
),
score: -201,
ply: 68,
result: 0,
},
TrainingDataEntry {
pos: Position::from_fen("1q5b/1r5k/4p2p/1b2P1pN/2Pp4/6PP/1n4B1/1Q2B1K1 b - - 0 35")
.unwrap(),
mv: Move::new(
Square::new(27),
Square::new(19),
MoveType::Normal,
Piece::none(),
),
score: 254,
ply: 69,
result: 0,
},
TrainingDataEntry {
pos: Position::from_fen(
"1q5b/1r5k/4p2p/1b2P1pN/2P5/3p2PP/1n4B1/1Q2B1K1 w - - 0 36",
)
.unwrap(),
mv: Move::new(
Square::new(14),
Square::new(49),
MoveType::Normal,
Piece::none(),
),
score: -220,
ply: 70,
result: 0,
},
];
assert_eq!(entries, expected);
}
#[test]
fn test_reader_big_score_diff() {
let cursor: Cursor<Vec<u8>> = Cursor::new(Vec::from([
66, 73, 78, 80, 37, 0, 0, 0, 130, 130, 144, 210, 8, 192, 70, 82, 72, 58, 64, 0, 81, 16,
18, 113, 155, 5, 0, 0, 0, 0, 0, 0, 10, 104, 249, 253, 0, 68, 0, 0, 0, 1, 29, 83, 79,
]));
let mut reader = CompressedTrainingDataEntryReader::new(cursor).unwrap();
let mut entries: Vec<TrainingDataEntry> = Vec::new();
while reader.has_next() {
let entry = reader.next();
entries.push(entry);
}
let expected = vec![
TrainingDataEntry {
pos: Position::from_fen("1q5b/1r5k/4p2p/1b2P1pN/3p4/6PP/1nP3B1/1Q2B1K1 w - - 0 35")
.unwrap(),
mv: Move::new(
Square::new(10),
Square::new(26),
MoveType::Normal,
Piece::none(),
),
score: -31999,
ply: 68,
result: 0,
},
TrainingDataEntry {
pos: Position::from_fen("1q5b/1r5k/4p2p/1b2P1pN/2Pp4/6PP/1n4B1/1Q2B1K1 b - - 0 35")
.unwrap(),
mv: Move::new(
Square::new(27),
Square::new(19),
MoveType::Normal,
Piece::none(),
),
score: -1500,
ply: 69,
result: 0,
},
];
assert_eq!(entries, expected);
}
#[test]
fn test_reader_from_bytes() {
let file = std::fs::read("./test/ep1.binpack").unwrap();
let mut reader = CompressedTrainingDataEntryReader::from_bytes(file).unwrap();
let mut num_entries = 0;
while reader.has_next() {
let _ = reader.next();
num_entries += 1;
}
assert_eq!(num_entries, 3);
}
#[test]
fn test_chunk_read_and_parse() {
let first_chunk: Vec<u8> = vec![
98, 121, 192, 21, 24, 76, 241, 100, 100, 106, 0, 4, 8, 48, 2, 17, 17, 145, 19, 117,
247, 0, 0, 0, 61, 232, 0, 253, 0, 39, 0, 2, 0, 0,
];
let second_chunk: Vec<u8> = vec![
98, 121, 192, 21, 24, 76, 241, 100, 100, 106, 0, 4, 8, 48, 2, 17, 17, 145, 19, 117,
247, 0, 0, 0, 61, 232, 0, 253, 0, 39, 0, 2, 0, 0,
];
let mut file = Vec::new();
file.extend_from_slice(b"BINP");
file.extend_from_slice(&(first_chunk.len() as u32).to_le_bytes());
file.extend_from_slice(&first_chunk);
file.extend_from_slice(b"BINP");
file.extend_from_slice(&(second_chunk.len() as u32).to_le_bytes());
file.extend_from_slice(&second_chunk);
let mut reader = CompressedTrainingDataEntryReader::from_bytes(file).unwrap();
let mut chunk = Vec::new();
assert!(reader.read_next_chunk_into(&mut chunk).unwrap());
assert_eq!(chunk, second_chunk);
let entries = parse_chunk(&chunk);
assert_eq!(entries.len(), 1);
assert!(!reader.read_next_chunk_into(&mut chunk).unwrap());
}
#[test]
#[should_panic(expected = "index out of bounds: the len is 0 but the index is 0")]
fn test_reader_no_moves() {
let entry_bytes: [u8; 32] = [
98, 121, 192, 21, 24, 76, 241, 100, 100, 106, 0, 4, 8, 48, 2, 17, 17, 145, 19, 117,
247, 0, 0, 0, 61, 232, 0, 253, 0, 39, 0, 2,
];
let mut chunk = Vec::new();
chunk.extend_from_slice(&entry_bytes);
chunk.extend_from_slice(&1u16.to_be_bytes());
let mut file = Vec::new();
file.extend_from_slice(b"BINP");
file.extend_from_slice(&(chunk.len() as u32).to_le_bytes());
file.extend_from_slice(&chunk);
let cursor = Cursor::new(file);
let mut reader = CompressedTrainingDataEntryReader::new(cursor).unwrap();
let _ = reader.next();
let _ = reader.next();
}
}