use std::io::{self, Read, Seek, Write};
use std::path::Path;
use rayon::prelude::*;
use serde::{Serialize, Deserialize, de::DeserializeOwned};
use super::writer::ZipDocumentWriter;
use super::reader::ZipDocumentReader;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SnapshotShardMeta {
pub root_id: String,
pub num_shards: usize,
pub counts: Vec<usize>,
}
#[cfg_attr(feature = "dev-tracing", tracing::instrument(skip(zw, meta, get_shard_bytes), fields(
crate_name = "file",
shard_count = meta.shard_count,
zstd_level = zstd_level
)))]
pub fn write_snapshot_shards<W, F>(
zw: &mut ZipDocumentWriter<W>,
meta: &SnapshotShardMeta,
mut get_shard_bytes: F,
zstd_level: i32,
) -> io::Result<()>
where
W: Write + Seek,
F: FnMut(usize) -> io::Result<Vec<u8>>,
{
let meta_val = serde_json::to_value(meta).map_err(io::Error::other)?;
zw.add_json("snapshot/meta.json", &meta_val)?;
for i in 0..meta.num_shards {
let raw = get_shard_bytes(i)?;
let zst = zstd::stream::encode_all(&raw[..], zstd_level)
.map_err(io::Error::other)?;
let name = format!("snapshot/shard-{i:03}.bin.zst");
zw.add_stored(&name, &zst)?;
}
Ok(())
}
#[cfg_attr(
feature = "dev-tracing",
tracing::instrument(skip(zr), fields(crate_name = "file"))
)]
pub fn read_snapshot_shards<R: Read + Seek>(
zr: &mut ZipDocumentReader<R>
) -> io::Result<(SnapshotShardMeta, Vec<Vec<u8>>)> {
let meta_bytes = zr.read_all("snapshot/meta.json")?;
let meta: SnapshotShardMeta =
serde_json::from_slice(&meta_bytes).map_err(io::Error::other)?;
let mut compressed: Vec<Vec<u8>> = Vec::with_capacity(meta.num_shards);
for i in 0..meta.num_shards {
let name = format!("snapshot/shard-{i:03}.bin.zst");
let zst = zr.read_all(&name)?;
compressed.push(zst);
}
let shards: Vec<Vec<u8>> = compressed
.into_par_iter()
.map(|zst| zstd::stream::decode_all(&zst[..]).map_err(io::Error::other))
.collect::<Result<Vec<_>, _>>()?;
Ok((meta, shards))
}
pub fn read_and_decode_snapshot_shards<R: Read + Seek, T: DeserializeOwned>(
zr: &mut ZipDocumentReader<R>
) -> io::Result<(SnapshotShardMeta, Vec<T>)> {
let (meta, shards_raw) = read_snapshot_shards(zr)?;
let mut out: Vec<T> = Vec::with_capacity(shards_raw.len());
for raw in shards_raw.iter() {
let (val, _): (T, _) =
bincode::serde::decode_from_slice(raw, bincode::config::standard())
.map_err(io::Error::other)?;
out.push(val);
}
Ok((meta, out))
}
pub fn for_each_snapshot_shard_raw<R: Read + Seek, F>(
zr: &mut ZipDocumentReader<R>,
mut on_shard: F,
) -> io::Result<SnapshotShardMeta>
where
F: FnMut(usize, Vec<u8>) -> io::Result<()>,
{
let meta_bytes = zr.read_all("snapshot/meta.json")?;
let meta: SnapshotShardMeta =
serde_json::from_slice(&meta_bytes).map_err(io::Error::other)?;
for i in 0..meta.num_shards {
let name = format!("snapshot/shard-{i:03}.bin.zst");
let zst = zr.read_all(&name)?;
let raw =
zstd::stream::decode_all(&zst[..]).map_err(io::Error::other)?;
on_shard(i, raw)?;
}
Ok(meta)
}
#[cfg_attr(feature = "dev-tracing", tracing::instrument(skip(path, meta_json, schema_xml, shard_meta, get_shard_bytes), fields(
crate_name = "file",
file_path = %path.as_ref().display(),
shard_count = shard_meta.counts,
schema_size = schema_xml.len(),
zstd_level = zstd_level
)))]
pub fn export_zip_with_shards<P, F>(
path: P,
meta_json: &serde_json::Value,
schema_xml: &[u8],
shard_meta: &SnapshotShardMeta,
get_shard_bytes: F,
zstd_level: i32,
) -> io::Result<()>
where
P: AsRef<Path>,
F: FnMut(usize) -> io::Result<Vec<u8>>,
{
let file = std::fs::File::create(path)?;
let mut zw = ZipDocumentWriter::new(file)?;
zw.add_json("meta.json", meta_json)?;
zw.add_deflated("schema.xml", schema_xml)?;
write_snapshot_shards(&mut zw, shard_meta, get_shard_bytes, zstd_level)?;
let _ = zw.finalize()?;
Ok(())
}
#[cfg_attr(feature = "dev-tracing", tracing::instrument(skip(path), fields(
crate_name = "file",
file_path = %path.as_ref().display()
)))]
pub fn import_zip_with_shards<P, T>(
path: P
) -> io::Result<(serde_json::Value, Vec<u8>, SnapshotShardMeta, Vec<T>)>
where
P: AsRef<Path>,
T: DeserializeOwned,
{
let file = std::fs::File::open(path)?;
let mut zr = ZipDocumentReader::new(file)?;
let meta_json = zr.read_all("meta.json")?;
let meta_val: serde_json::Value =
serde_json::from_slice(&meta_json).map_err(io::Error::other)?;
let schema_xml = zr.read_all("schema.xml")?;
let (shard_meta, decoded) =
read_and_decode_snapshot_shards::<_, T>(&mut zr)?;
Ok((meta_val, schema_xml, shard_meta, decoded))
}