use std::io::SeekFrom;
use axum::body::Bytes;
use futures_util::Stream;
use tokio::fs;
use tokio::io::{AsyncReadExt, AsyncSeekExt, AsyncWriteExt};
use super::crypt::{self, Keyring, ObjectKey};
use crate::error::Error;
const MAGIC: &[u8; 4] = b"LFZ1";
const V_PLAIN: u8 = 1;
const V_SEALED: u8 = 2;
const HEADER: u64 = 32;
const SEALED_HEADER: u64 = 64;
fn header_len(version: u8) -> u64 {
match version {
V_SEALED => SEALED_HEADER,
_ => HEADER,
}
}
const STORED: u32 = 1 << 31;
pub const FRAME: u64 = 4 * 1024 * 1024;
pub struct Framed {
file: fs::File,
plaintext: u64,
frame: u64,
frames: u32,
offsets: Vec<u64>,
stored: Vec<bool>,
key: Option<ObjectKey>,
oid: String,
}
pub struct Writer {
file: fs::File,
level: Option<i32>,
key: Option<ObjectKey>,
oid: String,
lengths: Vec<u32>,
plaintext: u64,
pending: Vec<u8>,
}
impl Writer {
pub async fn open(
mut file: fs::File,
level: Option<i32>,
key: Option<ObjectKey>,
oid: &str,
) -> Result<Self, Error> {
let version = if key.is_some() { V_SEALED } else { V_PLAIN };
file.write_all(&vec![0u8; header_len(version) as usize])
.await?;
Ok(Self {
file,
level,
key,
oid: oid.to_owned(),
lengths: Vec::new(),
plaintext: 0,
pending: Vec::with_capacity(FRAME as usize),
})
}
pub async fn push(&mut self, chunk: &[u8]) -> Result<(), Error> {
self.plaintext += chunk.len() as u64;
self.pending.extend_from_slice(chunk);
while self.pending.len() > FRAME as usize {
let rest = self.pending.split_off(FRAME as usize);
let frame = std::mem::replace(&mut self.pending, rest);
self.flush(&frame, false).await?;
}
Ok(())
}
pub async fn finish(mut self) -> Result<(), Error> {
let frame = std::mem::take(&mut self.pending);
if !frame.is_empty() {
self.flush(&frame, true).await?;
}
for length in &self.lengths {
self.file.write_all(&length.to_le_bytes()).await?;
}
let header = header_len(self.version());
let index = header
+ self
.lengths
.iter()
.map(|length| (*length & !STORED) as u64)
.sum::<u64>();
self.file.seek(SeekFrom::Start(0)).await?;
self.file.write_all(&self.header(index)).await?;
self.file.flush().await?;
self.file.sync_all().await?;
Ok(())
}
fn version(&self) -> u8 {
if self.key.is_some() {
V_SEALED
} else {
V_PLAIN
}
}
fn header(&self, index: u64) -> Vec<u8> {
let version = self.version();
let mut header = vec![0u8; header_len(version) as usize];
header[0..4].copy_from_slice(MAGIC);
header[4] = version;
header[8..16].copy_from_slice(&self.plaintext.to_le_bytes());
header[16..20].copy_from_slice(&(FRAME as u32).to_le_bytes());
header[20..24].copy_from_slice(&(self.lengths.len() as u32).to_le_bytes());
header[24..32].copy_from_slice(&index.to_le_bytes());
if let Some(key) = &self.key {
header[32..36].copy_from_slice(&key.id());
header[36..52].copy_from_slice(&key.salt());
}
header
}
async fn flush(&mut self, frame: &[u8], last: bool) -> Result<(), Error> {
let (body, stored) = match self.level {
Some(level) => {
let compressed =
zstd::bulk::compress(frame, level).map_err(std::io::Error::other)?;
if compressed.len() >= frame.len() - frame.len() / 20 {
(frame.to_vec(), true)
} else {
(compressed, false)
}
}
None => (frame.to_vec(), true),
};
let index = self.lengths.len() as u32;
let body = match &self.key {
Some(key) => key.seal(index, last, &self.oid, &body)?,
None => body,
};
self.file.write_all(&body).await?;
self.lengths
.push(body.len() as u32 | if stored { STORED } else { 0 });
Ok(())
}
}
const FRAME_MIN: u64 = 64 * 1024;
const FRAME_MAX: u64 = 16 * 1024 * 1024;
fn plausible(plaintext: u64, frame: u64, frames: u64) -> bool {
if !(FRAME_MIN..=FRAME_MAX).contains(&frame) {
return false;
}
frames == plaintext.div_ceil(frame)
}
impl Framed {
pub async fn open(
mut file: fs::File,
on_disk: u64,
keys: Option<&Keyring>,
oid: &str,
) -> Result<Option<Self>, Error> {
if on_disk < header_len(V_PLAIN) {
return Ok(None);
}
let mut prefix = [0u8; 32];
file.read_exact(&mut prefix).await?;
if &prefix[0..4] != MAGIC || !matches!(prefix[4], V_PLAIN | V_SEALED) {
file.seek(SeekFrom::Start(0)).await?;
return Ok(None);
}
let version = prefix[4];
let header = header_len(version);
let plaintext = u64::from_le_bytes(prefix[8..16].try_into().expect("eight bytes"));
let frame = u32::from_le_bytes(prefix[16..20].try_into().expect("four bytes")) as u64;
let frames = u32::from_le_bytes(prefix[20..24].try_into().expect("four bytes")) as u64;
let index = u64::from_le_bytes(prefix[24..32].try_into().expect("eight bytes"));
let indexed = frames.saturating_mul(4);
if !plausible(plaintext, frame, frames)
|| index < header
|| index.saturating_add(indexed) != on_disk
{
file.seek(SeekFrom::Start(0)).await?;
return Ok(None);
}
let key = match version {
V_SEALED => Some(sealed_key(&mut file, keys).await?),
_ => None,
};
file.seek(SeekFrom::Start(index)).await?;
let mut lengths = vec![0u8; indexed as usize];
file.read_exact(&mut lengths).await?;
let mut offsets = Vec::with_capacity(frames as usize + 1);
let mut stored = Vec::with_capacity(frames as usize);
let mut at = header;
offsets.push(at);
for length in lengths.chunks_exact(4) {
let entry = u32::from_le_bytes(length.try_into().expect("four bytes"));
at += (entry & !STORED) as u64;
stored.push(entry & STORED != 0);
offsets.push(at);
}
if at != index {
file.seek(SeekFrom::Start(0)).await?;
return Ok(None);
}
Ok(Some(Self {
file,
plaintext,
frame,
frames: frames as u32,
offsets,
stored,
key,
oid: oid.to_owned(),
}))
}
pub fn plaintext(&self) -> u64 {
self.plaintext
}
pub fn stream(self, start: u64, length: u64) -> impl Stream<Item = Result<Bytes, Error>> {
futures_util::stream::try_unfold(
(self, start, length),
|(mut framed, at, wanted)| async move {
if wanted == 0 {
return Ok(None);
}
let index = (at / framed.frame) as usize;
let Some(bounds) = framed.offsets.get(index..index + 2).map(<[u64]>::to_vec) else {
return Ok(None);
};
framed.file.seek(SeekFrom::Start(bounds[0])).await?;
let mut body = vec![0u8; (bounds[1] - bounds[0]) as usize];
framed.file.read_exact(&mut body).await?;
let plain = framed.decode(index, body)?;
let from = (at % framed.frame) as usize;
let take = wanted.min(plain.len().saturating_sub(from) as u64) as usize;
let bytes = Bytes::copy_from_slice(&plain[from..from + take]);
Ok(Some((
bytes,
(framed, at + take as u64, wanted - take as u64),
)))
},
)
}
fn decode(&self, index: usize, body: Vec<u8>) -> Result<Vec<u8>, Error> {
let body = match &self.key {
Some(key) => key.open(
index as u32,
index as u32 + 1 == self.frames,
&self.oid,
&body,
)?,
None => body,
};
if self.stored.get(index).copied().unwrap_or_default() {
return Ok(body);
}
zstd::bulk::decompress(&body, self.frame as usize)
.map_err(|error| Error::Storage(std::io::Error::other(error)))
}
}
async fn sealed_key(file: &mut fs::File, keys: Option<&Keyring>) -> Result<ObjectKey, Error> {
let keys = keys.ok_or(Error::NotDecryptable)?;
let mut rest = [0u8; 32];
file.read_exact(&mut rest).await?;
let mut id = [0u8; crypt::ID];
id.copy_from_slice(&rest[0..crypt::ID]);
let mut salt = [0u8; crypt::SALT];
salt.copy_from_slice(&rest[4..4 + crypt::SALT]);
keys.reading(id, salt)
}
#[cfg(test)]
mod tests;