use std::{
fs::{self, File, OpenOptions},
io::{Read, Write},
mem,
path::{Path, PathBuf},
};
pub use whasher::{StreamHasher, compute_checksum, compute_checksum_with_seed};
use crate::{
error::{Error, Result},
stub::RANGE_INDEX_STUB_SIZE,
};
pub const KEY_LEN_BYTES: usize = 4;
pub const FILE_LEN_BYTES: usize = 8;
pub const CHECKSUM_BYTES: usize = 8;
pub const STUB_LEN_BYTES: usize = 4;
pub const MAX_KEY_LEN_BYTES: usize = 64 * 1024 * 1024;
pub const MAX_FILE_LEN_BYTES: u64 = 64 * 1024 * 1024 * 1024;
pub const MIN_CHUNK_SIZE: usize = CHECKSUM_BYTES + STUB_LEN_BYTES + RANGE_INDEX_STUB_SIZE;
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
enum SerializerPhase {
KeyHeader,
KeyData,
FileHeader,
FileData,
Trailer,
Done,
}
pub struct RangeIndexChunkedSerializer {
key_bytes: Vec<u8>,
stub_bytes: Vec<u8>,
total_file_bytes: u64,
file_bytes_emitted: u64,
key_bytes_emitted: usize,
hasher: StreamHasher,
phase: SerializerPhase,
unprocessed_file_data: Vec<u8>,
unprocessed_offset: usize,
}
impl RangeIndexChunkedSerializer {
pub fn new(key: &[u8], stub: &[u8], total_file_bytes: u64) -> Self {
Self {
key_bytes: key.to_vec(),
stub_bytes: stub.to_vec(),
total_file_bytes,
file_bytes_emitted: 0,
key_bytes_emitted: 0,
hasher: StreamHasher::default(),
phase: SerializerPhase::KeyHeader,
unprocessed_file_data: Vec::new(),
unprocessed_offset: 0,
}
}
pub fn new_with_checksum(key: &[u8], stub: &[u8], total_file_bytes: u64, _checksum: u64) -> Self {
Self::new(key, stub, total_file_bytes)
}
#[inline]
pub fn total_file_bytes(&self) -> u64 {
self.total_file_bytes
}
#[inline]
pub fn is_complete(&self) -> bool {
self.phase == SerializerPhase::Done
}
#[inline]
pub fn needs_file_data(&self) -> bool {
self.phase == SerializerPhase::FileData
&& self.file_bytes_emitted < self.total_file_bytes
&& self.unprocessed_offset >= self.unprocessed_file_data.len()
}
#[inline]
pub fn file_data_remaining(&self) -> u64 {
self
.total_file_bytes
.saturating_sub(self.file_bytes_emitted)
}
pub fn supply_file_data(&mut self, data: &[u8]) {
self.unprocessed_file_data.clear();
self.unprocessed_file_data.extend_from_slice(data);
self.unprocessed_offset = 0;
}
pub fn move_next(&mut self, dest: &mut [u8]) -> Result<usize> {
if self.phase == SerializerPhase::Done {
return Err(Error::InvalidArgument(
"Serializer has already completed".into(),
));
}
let dest_capacity = dest.len();
let mut written = 0;
if self.phase == SerializerPhase::KeyHeader {
if dest_capacity - written < KEY_LEN_BYTES {
return Ok(written);
}
let key_len = (self.key_bytes.len() as u32).to_le_bytes();
dest[written..written + KEY_LEN_BYTES].copy_from_slice(&key_len);
written += KEY_LEN_BYTES;
self.phase = SerializerPhase::KeyData;
}
if self.phase == SerializerPhase::KeyData {
let remain_key = self.key_bytes.len() - self.key_bytes_emitted;
let avail_dest = dest_capacity - written;
let to_copy = remain_key.min(avail_dest);
dest[written..written + to_copy]
.copy_from_slice(&self.key_bytes[self.key_bytes_emitted..self.key_bytes_emitted + to_copy]);
written += to_copy;
self.key_bytes_emitted += to_copy;
if self.key_bytes_emitted < self.key_bytes.len() {
return Ok(written);
}
self.phase = SerializerPhase::FileHeader;
}
if self.phase == SerializerPhase::FileHeader {
if dest_capacity - written < FILE_LEN_BYTES {
return Ok(written);
}
dest[written..written + FILE_LEN_BYTES].copy_from_slice(&self.total_file_bytes.to_le_bytes());
written += FILE_LEN_BYTES;
self.phase = SerializerPhase::FileData;
}
if self.phase == SerializerPhase::FileData {
if self.file_bytes_emitted < self.total_file_bytes {
let avail_dest = dest_capacity - written;
if avail_dest == 0 {
return Ok(written);
}
let max_copy = ((self.total_file_bytes - self.file_bytes_emitted) as usize).min(avail_dest);
let unproc_avail = self.unprocessed_file_data.len() - self.unprocessed_offset;
let to_copy = max_copy.min(unproc_avail);
if to_copy == 0 {
return Ok(written);
}
let src_slice =
&self.unprocessed_file_data[self.unprocessed_offset..self.unprocessed_offset + to_copy];
dest[written..written + to_copy].copy_from_slice(src_slice);
self.hasher.write(src_slice);
written += to_copy;
self.unprocessed_offset += to_copy;
self.file_bytes_emitted += to_copy as u64;
}
if self.file_bytes_emitted >= self.total_file_bytes {
self.phase = SerializerPhase::Trailer;
}
}
if self.phase == SerializerPhase::Trailer {
let trailer_len = CHECKSUM_BYTES + STUB_LEN_BYTES + self.stub_bytes.len();
if dest_capacity - written < trailer_len {
return Ok(written);
}
let actual_checksum = self.hasher.finish();
dest[written..written + CHECKSUM_BYTES].copy_from_slice(&actual_checksum.to_le_bytes());
written += CHECKSUM_BYTES;
let stub_len = (self.stub_bytes.len() as u32).to_le_bytes();
dest[written..written + STUB_LEN_BYTES].copy_from_slice(&stub_len);
written += STUB_LEN_BYTES;
dest[written..written + self.stub_bytes.len()].copy_from_slice(&self.stub_bytes);
written += self.stub_bytes.len();
self.phase = SerializerPhase::Done;
}
Ok(written)
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
enum DeserializerState {
WaitingForKeyHeader,
ReceivingKeyData,
WaitingForFileHeader,
ReceivingFileData,
WaitingForTrailer,
Complete,
Error,
Disposed,
}
pub struct RangeIndexChunkedDeserializer {
temp_path: PathBuf,
file: Option<File>,
state: DeserializerState,
key: Vec<u8>,
key_bytes_received: usize,
total_file_bytes: u64,
file_bytes_remaining: u64,
hasher: StreamHasher,
finalizer_stub: Vec<u8>,
}
impl RangeIndexChunkedDeserializer {
pub fn new(temp_path: impl Into<PathBuf>) -> Result<Self> {
let temp_path = temp_path.into();
Ok(Self {
temp_path,
file: None,
state: DeserializerState::WaitingForKeyHeader,
key: Vec::new(),
key_bytes_received: 0,
total_file_bytes: 0,
file_bytes_remaining: 0,
hasher: StreamHasher::default(),
finalizer_stub: Vec::new(),
})
}
#[inline]
pub fn is_complete(&self) -> bool {
self.state == DeserializerState::Complete
}
#[inline]
pub fn has_error(&self) -> bool {
self.state == DeserializerState::Error
}
#[inline]
pub fn key(&self) -> &[u8] {
&self.key
}
#[inline]
pub fn stub(&self) -> &[u8] {
&self.finalizer_stub
}
#[inline]
pub fn temp_path(&self) -> &Path {
&self.temp_path
}
pub fn process_chunk(&mut self, mut data: &[u8]) -> Result<bool> {
loop {
match self.state {
DeserializerState::Error | DeserializerState::Complete | DeserializerState::Disposed => {
return Ok(false);
}
DeserializerState::WaitingForKeyHeader => {
if data.is_empty() {
return Ok(true);
}
if data.len() < KEY_LEN_BYTES {
self.state = DeserializerState::Error;
return Ok(false);
}
let key_len_i32 = i32::from_le_bytes(data[..KEY_LEN_BYTES].try_into().unwrap());
data = &data[KEY_LEN_BYTES..];
if key_len_i32 <= 0 || (key_len_i32 as usize) > MAX_KEY_LEN_BYTES {
self.state = DeserializerState::Error;
return Ok(false);
}
let key_len = key_len_i32 as usize;
self.key = vec![0u8; key_len];
self.key_bytes_received = 0;
self.state = DeserializerState::ReceivingKeyData;
}
DeserializerState::ReceivingKeyData => {
let needed = self.key.len() - self.key_bytes_received;
let n = needed.min(data.len());
self.key[self.key_bytes_received..self.key_bytes_received + n]
.copy_from_slice(&data[..n]);
self.key_bytes_received += n;
data = &data[n..];
if self.key_bytes_received < self.key.len() {
return Ok(true);
}
self.state = DeserializerState::WaitingForFileHeader;
}
DeserializerState::WaitingForFileHeader => {
if data.is_empty() {
return Ok(true);
}
if data.len() < FILE_LEN_BYTES {
self.state = DeserializerState::Error;
return Ok(false);
}
let file_len_i64 = i64::from_le_bytes(data[..FILE_LEN_BYTES].try_into().unwrap());
data = &data[FILE_LEN_BYTES..];
if file_len_i64 <= 0 || (file_len_i64 as u64) > MAX_FILE_LEN_BYTES {
self.state = DeserializerState::Error;
return Ok(false);
}
let file_len = file_len_i64 as u64;
self.total_file_bytes = file_len;
self.file_bytes_remaining = file_len;
self.state = DeserializerState::ReceivingFileData;
match OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(&self.temp_path)
{
Ok(f) => self.file = Some(f),
Err(_) => {
self.state = DeserializerState::Error;
return Ok(false);
}
}
}
DeserializerState::ReceivingFileData => {
if data.is_empty() {
return Ok(true);
}
if self.file_bytes_remaining > 0 {
let count = (data.len() as u64).min(self.file_bytes_remaining) as usize;
let file_part = &data[..count];
if let Some(ref mut f) = self.file {
if f.write_all(file_part).is_err() {
self.state = DeserializerState::Error;
return Ok(false);
}
} else {
self.state = DeserializerState::Error;
return Ok(false);
}
self.hasher.write(file_part);
self.file_bytes_remaining -= count as u64;
data = &data[count..];
}
if self.file_bytes_remaining == 0 {
if let Some(mut f) = self.file.take()
&& (f.flush().is_err() || f.sync_all().is_err())
{
self.state = DeserializerState::Error;
return Ok(false);
}
self.state = DeserializerState::WaitingForTrailer;
} else {
return Ok(true);
}
}
DeserializerState::WaitingForTrailer => {
if data.is_empty() {
return Ok(true);
}
let min_trailer_len = CHECKSUM_BYTES + STUB_LEN_BYTES;
if data.len() < min_trailer_len {
self.state = DeserializerState::Error;
return Ok(false);
}
let received_hash = u64::from_le_bytes(data[..CHECKSUM_BYTES].try_into().unwrap());
data = &data[CHECKSUM_BYTES..];
let stub_len = i32::from_le_bytes(data[..STUB_LEN_BYTES].try_into().unwrap());
data = &data[STUB_LEN_BYTES..];
if stub_len != RANGE_INDEX_STUB_SIZE as i32 {
self.state = DeserializerState::Error;
return Ok(false);
}
if data.len() != RANGE_INDEX_STUB_SIZE {
self.state = DeserializerState::Error;
return Ok(false);
}
let calculated_checksum = self.hasher.finish();
if received_hash != calculated_checksum {
self.state = DeserializerState::Error;
return Ok(false);
}
self.finalizer_stub = data.to_vec();
self.state = DeserializerState::Complete;
return Ok(true);
}
}
}
}
pub fn dispose(&mut self) {
if self.state == DeserializerState::Disposed {
return;
}
self.state = DeserializerState::Disposed;
self.file.take();
if self.temp_path.exists() {
let _ = fs::remove_file(&self.temp_path);
}
}
pub fn take_temp_path(mut self) -> PathBuf {
self.state = DeserializerState::Disposed;
self.file.take();
mem::take(&mut self.temp_path)
}
}
impl Drop for RangeIndexChunkedDeserializer {
fn drop(&mut self) {
self.dispose();
}
}
pub struct RangeIndexMigrationReader<R: Read> {
serializer: RangeIndexChunkedSerializer,
reader: Option<R>,
temp_file_path: Option<PathBuf>,
read_buffer: Vec<u8>,
disposed: bool,
}
impl<R: Read> RangeIndexMigrationReader<R> {
pub fn new(
serializer: RangeIndexChunkedSerializer,
reader: R,
temp_file_path: Option<PathBuf>,
read_buffer_size: usize,
) -> Result<Self> {
if read_buffer_size == 0 {
return Err(Error::InvalidArgument(
"read_buffer_size must be positive".into(),
));
}
Ok(Self {
serializer,
reader: Some(reader),
temp_file_path,
read_buffer: vec![0u8; read_buffer_size],
disposed: false,
})
}
#[inline]
pub fn is_complete(&self) -> bool {
self.serializer.is_complete()
}
#[inline]
pub fn total_file_bytes(&self) -> u64 {
self.serializer.total_file_bytes()
}
pub fn read_next_chunk(&mut self, mut destination: &mut [u8]) -> Result<usize> {
if destination.len() < MIN_CHUNK_SIZE {
return Err(Error::InvalidArgument(format!(
"destination must be at least {MIN_CHUNK_SIZE} bytes (the trailer size) so the stream can complete"
)));
}
let initial_len = destination.len();
let reader = match self.reader.as_mut() {
Some(r) => r,
None => return Err(Error::InvalidArgument("Reader already disposed".into())),
};
while !self.serializer.is_complete() && !destination.is_empty() {
if self.serializer.needs_file_data() {
let max_read =
(self.read_buffer.len() as u64).min(self.serializer.file_data_remaining()) as usize;
let bytes_read = reader.read(&mut self.read_buffer[..max_read])?;
if bytes_read == 0 && self.serializer.file_data_remaining() > 0 {
return Err(Error::Corrupted(format!(
"RangeIndex file truncated: {} bytes remaining",
self.serializer.file_data_remaining()
)));
}
self
.serializer
.supply_file_data(&self.read_buffer[..bytes_read]);
}
let written = self.serializer.move_next(destination)?;
if written == 0 {
break;
}
destination = &mut destination[written..];
}
Ok(initial_len - destination.len())
}
pub fn dispose(&mut self) {
if self.disposed {
return;
}
self.disposed = true;
self.reader.take();
if let Some(ref path) = self.temp_file_path
&& path.exists()
{
let _ = fs::remove_file(path);
}
}
}
impl<R: Read> Drop for RangeIndexMigrationReader<R> {
fn drop(&mut self) {
self.dispose();
}
}