use std::fs::{self, File, OpenOptions};
use std::io::{self, BufWriter, Write};
use std::path::{Path, PathBuf};
use flate2::Compression;
use flate2::write::GzEncoder;
use zstd::stream::write::Encoder as ZstdEncoder;
use crate::cli::RecordCompression;
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum Codec {
Plain,
Zstd,
Gzip,
}
impl Codec {
pub fn detect(path: &Path, override_: Option<RecordCompression>) -> Self {
if let Some(r) = override_ {
return match r {
RecordCompression::Zstd => Codec::Zstd,
RecordCompression::Gzip => Codec::Gzip,
RecordCompression::None => Codec::Plain,
};
}
match path
.extension()
.and_then(|e| e.to_str())
.map(|e| e.to_ascii_lowercase())
.as_deref()
{
Some("zst") => Codec::Zstd,
Some("gz") => Codec::Gzip,
_ => Codec::Plain,
}
}
}
enum ActiveSegment {
Plain(BufWriter<File>),
Zstd(ZstdEncoder<'static, BufWriter<File>>),
Gzip(GzEncoder<BufWriter<File>>),
}
impl ActiveSegment {
fn open(path: &Path, codec: Codec) -> io::Result<Self> {
let file = open_secure(path)?;
let buf = BufWriter::with_capacity(64 * 1024, file);
match codec {
Codec::Plain => Ok(Self::Plain(buf)),
Codec::Zstd => {
let mut enc = ZstdEncoder::new(buf, 3)?;
enc.include_checksum(true)?;
Ok(Self::Zstd(enc))
}
Codec::Gzip => Ok(Self::Gzip(GzEncoder::new(buf, Compression::default()))),
}
}
fn flush(&mut self) -> io::Result<()> {
match self {
Self::Plain(w) => w.flush(),
Self::Zstd(w) => w.flush(),
Self::Gzip(w) => w.flush(),
}
}
fn finish(self) -> io::Result<()> {
match self {
Self::Plain(mut w) => {
w.flush()?;
let mut inner = w.into_inner().map_err(|e| e.into_error())?;
inner.flush()?;
Ok(())
}
Self::Zstd(enc) => {
let buf = enc.finish()?;
let mut inner = buf.into_inner().map_err(|e| e.into_error())?;
inner.flush()?;
Ok(())
}
Self::Gzip(enc) => {
let buf = enc.finish()?;
let mut inner = buf.into_inner().map_err(|e| e.into_error())?;
inner.flush()?;
Ok(())
}
}
}
}
impl Write for ActiveSegment {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
match self {
Self::Plain(w) => w.write(buf),
Self::Zstd(w) => w.write(buf),
Self::Gzip(w) => w.write(buf),
}
}
fn flush(&mut self) -> io::Result<()> {
ActiveSegment::flush(self)
}
}
pub struct RotatingWriter {
base: PathBuf,
codec: Codec,
max_size: u64,
max_files: u32,
active: Option<ActiveSegment>,
active_bytes: u64,
next_rollover_index: u32,
rotated_segments: Vec<PathBuf>,
}
impl RotatingWriter {
pub fn new(
base: impl Into<PathBuf>,
codec: Codec,
max_size: u64,
max_files: u32,
) -> io::Result<Self> {
let base = base.into();
if let Some(parent) = base.parent()
&& !parent.as_os_str().is_empty()
{
fs::create_dir_all(parent)?;
}
let active = ActiveSegment::open(&base, codec)?;
Ok(Self {
base,
codec,
max_size,
max_files: max_files.max(1),
active: Some(active),
active_bytes: 0,
next_rollover_index: 1,
rotated_segments: Vec::new(),
})
}
pub fn write_line(&mut self, line: &[u8]) -> io::Result<()> {
let segment = self
.active
.as_mut()
.expect("RotatingWriter::active was taken; new segment not opened");
segment.write_all(line)?;
self.active_bytes += line.len() as u64;
if self.max_size > 0 && self.active_bytes >= self.max_size {
self.rollover()?;
}
Ok(())
}
fn rollover(&mut self) -> io::Result<()> {
let active = self
.active
.take()
.expect("RotatingWriter::active was already taken");
active.finish()?;
if self.max_files <= 1 {
self.active = Some(ActiveSegment::open(&self.base, self.codec)?);
self.active_bytes = 0;
self.next_rollover_index = self.next_rollover_index.saturating_add(1);
return Ok(());
}
let rolled_path = rotated_path(&self.base, self.next_rollover_index);
match fs::symlink_metadata(&rolled_path) {
Ok(m) if m.file_type().is_symlink() => {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
format!(
"refusing to rotate onto symlink at `{}`",
rolled_path.display()
),
));
}
Ok(_) => {
let _ = fs::remove_file(&rolled_path);
}
Err(e) if e.kind() == io::ErrorKind::NotFound => {
}
Err(e) => return Err(e),
}
fs::rename(&self.base, &rolled_path)?;
self.rotated_segments.push(rolled_path);
self.next_rollover_index = self.next_rollover_index.saturating_add(1);
while self.rotated_segments.len() as u32 >= self.max_files {
let oldest = self.rotated_segments.remove(0);
let _ = fs::remove_file(&oldest);
}
self.active = Some(ActiveSegment::open(&self.base, self.codec)?);
self.active_bytes = 0;
Ok(())
}
#[allow(dead_code)]
pub fn flush(&mut self) -> io::Result<()> {
if let Some(active) = self.active.as_mut() {
active.flush()?;
}
Ok(())
}
pub fn finish(mut self) -> io::Result<()> {
if let Some(active) = self.active.take() {
active.finish()?;
}
Ok(())
}
pub fn active_bytes(&self) -> u64 {
self.active_bytes
}
}
impl Drop for RotatingWriter {
fn drop(&mut self) {
if let Some(active) = self.active.take() {
let _ = active.finish();
}
}
}
fn rotated_path(base: &Path, index: u32) -> PathBuf {
let dir = base.parent().unwrap_or_else(|| Path::new(""));
let name = base
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("record");
let (stem, suffix) = match name.find('.') {
Some(pos) => (&name[..pos], &name[pos..]),
None => (name, ""),
};
let new_name = format!("{stem}.{index:04}{suffix}");
dir.join(new_name)
}
fn open_secure(path: &Path) -> io::Result<File> {
#[cfg(unix)]
{
use std::os::unix::fs::OpenOptionsExt;
OpenOptions::new()
.write(true)
.create_new(true)
.custom_flags(libc::O_NOFOLLOW)
.mode(0o600)
.open(path)
}
#[cfg(windows)]
{
use std::os::windows::fs::OpenOptionsExt;
OpenOptions::new()
.write(true)
.create_new(true)
.share_mode(0)
.open(path)
}
#[cfg(not(any(unix, windows)))]
{
OpenOptions::new().write(true).create_new(true).open(path)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Read;
#[test]
fn codec_detect_from_extension() {
assert_eq!(
Codec::detect(Path::new("a.ndjson"), None),
Codec::Plain,
"plain ndjson"
);
assert_eq!(
Codec::detect(Path::new("a.ndjson.zst"), None),
Codec::Zstd,
"zst wins"
);
assert_eq!(
Codec::detect(Path::new("a.ndjson.gz"), None),
Codec::Gzip,
"gz wins"
);
assert_eq!(
Codec::detect(Path::new("a.ndjson.zst"), Some(RecordCompression::None)),
Codec::Plain,
"explicit override beats extension"
);
}
#[test]
fn rotated_path_inserts_index_before_extension() {
assert_eq!(
rotated_path(Path::new("/tmp/out.ndjson.zst"), 3),
PathBuf::from("/tmp/out.0003.ndjson.zst"),
);
assert_eq!(
rotated_path(Path::new("out.ndjson"), 7),
PathBuf::from("out.0007.ndjson"),
);
}
#[test]
fn plain_writer_round_trip() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("rec.ndjson");
let mut w = RotatingWriter::new(&path, Codec::Plain, 0, 1).unwrap();
w.write_line(b"{\"a\":1}\n").unwrap();
w.write_line(b"{\"b\":2}\n").unwrap();
w.finish().unwrap();
let content = std::fs::read_to_string(&path).unwrap();
assert_eq!(content, "{\"a\":1}\n{\"b\":2}\n");
}
#[test]
fn gzip_writer_round_trip() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("rec.ndjson.gz");
let mut w = RotatingWriter::new(&path, Codec::Gzip, 0, 1).unwrap();
w.write_line(b"hello\n").unwrap();
w.finish().unwrap();
let f = std::fs::File::open(&path).unwrap();
let mut dec = flate2::read::GzDecoder::new(f);
let mut s = String::new();
dec.read_to_string(&mut s).unwrap();
assert_eq!(s, "hello\n");
}
#[test]
fn zstd_writer_round_trip() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("rec.ndjson.zst");
let mut w = RotatingWriter::new(&path, Codec::Zstd, 0, 1).unwrap();
w.write_line(b"hello\n").unwrap();
w.write_line(b"world\n").unwrap();
w.finish().unwrap();
let f = std::fs::File::open(&path).unwrap();
let mut dec = zstd::stream::read::Decoder::new(f).unwrap();
let mut s = String::new();
dec.read_to_string(&mut s).unwrap();
assert_eq!(s, "hello\nworld\n");
}
#[test]
fn rotation_evicts_oldest_beyond_max_files() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("rec.ndjson");
let mut w = RotatingWriter::new(&path, Codec::Plain, 3, 3).unwrap();
for i in 0..6 {
w.write_line(format!("{i}\n").as_bytes()).unwrap();
}
w.finish().unwrap();
let mut files: Vec<String> = std::fs::read_dir(dir.path())
.unwrap()
.filter_map(|e| e.ok())
.filter_map(|e| e.file_name().into_string().ok())
.filter(|n| n.starts_with("rec."))
.collect();
files.sort();
assert!(
files.len() <= 3,
"max_files not respected, found: {files:?}"
);
}
#[cfg(unix)]
#[test]
fn record_writer_refuses_preexisting_symlink() {
use std::os::unix::fs::symlink;
let dir = tempfile::tempdir().unwrap();
let target = dir.path().join("target");
std::fs::write(&target, "do-not-clobber\n").unwrap();
let record_path = dir.path().join("rec.ndjson");
symlink(&target, &record_path).unwrap();
let result = RotatingWriter::new(&record_path, Codec::Plain, 0, 1);
assert!(
result.is_err(),
"open_secure must refuse to follow a symlink"
);
let content = std::fs::read_to_string(&target).unwrap();
assert_eq!(content, "do-not-clobber\n");
}
#[cfg(unix)]
#[test]
fn record_writer_refuses_preexisting_regular_file() {
let dir = tempfile::tempdir().unwrap();
let record_path = dir.path().join("rec.ndjson");
std::fs::write(&record_path, "preexisting\n").unwrap();
let result = RotatingWriter::new(&record_path, Codec::Plain, 0, 1);
assert!(
result.is_err(),
"open_secure must refuse to clobber an existing regular file"
);
}
#[cfg(unix)]
#[test]
fn record_writer_sets_restrictive_mode() {
use std::os::unix::fs::PermissionsExt;
let dir = tempfile::tempdir().unwrap();
let record_path = dir.path().join("rec.ndjson");
let mut w = RotatingWriter::new(&record_path, Codec::Plain, 0, 1).unwrap();
w.write_line(b"{\"a\":1}\n").unwrap();
w.finish().unwrap();
let meta = std::fs::metadata(&record_path).unwrap();
let mode = meta.permissions().mode() & 0o777;
assert_eq!(mode, 0o600, "recording file must be 0o600, got {mode:o}");
}
#[cfg(unix)]
#[test]
fn record_writer_rollover_refuses_symlink_target() {
use std::os::unix::fs::symlink;
let dir = tempfile::tempdir().unwrap();
let record_path = dir.path().join("rec.ndjson");
let mut w = RotatingWriter::new(&record_path, Codec::Plain, 3, 3).unwrap();
let rolled = rotated_path(&record_path, 1);
let decoy = dir.path().join("decoy");
std::fs::write(&decoy, "preserve-me\n").unwrap();
symlink(&decoy, &rolled).unwrap();
let err = w.write_line(b"trigger\n").err();
assert!(err.is_some(), "rollover onto a symlink must raise an error");
let content = std::fs::read_to_string(&decoy).unwrap();
assert_eq!(content, "preserve-me\n");
}
}