use std::io::{self, Seek, Write};
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tap::Tap;
use tokio::sync::Mutex;
use tokio::task::JoinError;
pub struct BuilderExt<W: Write + Seek = OwnedOutput> {
tar: Arc<Mutex<BlowFuseOnDrop<W>>>,
path: PathBuf,
}
type OwnedOutput = Box<dyn WriteSeek + Send + 'static>;
type BorrowedOutput<'a> = Box<dyn WriteSeek + 'a>;
pub trait WriteSeek: Write + Seek {}
impl<T: Write + Seek> WriteSeek for T {}
struct BlowFuseOnDrop<W: Write + Seek> {
tar: Option<tar::Builder<FusedWriteSeek<W>>>,
enabled: Arc<AtomicBool>,
}
struct FusedWriteSeek<W> {
output: W,
enabled: Arc<AtomicBool>,
}
impl<W: Write + Seek> BlowFuseOnDrop<W> {
fn tar(&mut self) -> &mut tar::Builder<FusedWriteSeek<W>> {
self.tar.as_mut().unwrap()
}
}
impl<W: Write + Seek> Drop for BlowFuseOnDrop<W> {
fn drop(&mut self) {
self.enabled.store(false, Ordering::Release);
}
}
impl<W: Write> Write for FusedWriteSeek<W> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
if !self.enabled.load(Ordering::Acquire) {
return Err(io::Error::other("Using WriteBox after it is disabled"));
}
self.output.write(buf)
}
fn flush(&mut self) -> io::Result<()> {
self.output.flush()
}
}
impl<W: Seek> Seek for FusedWriteSeek<W> {
fn seek(&mut self, pos: io::SeekFrom) -> io::Result<u64> {
self.output.seek(pos)
}
}
impl BuilderExt<OwnedOutput> {
pub fn new_seekable_owned(output: impl Write + Seek + Send + 'static) -> Self {
Self::new(Box::new(output))
}
pub fn new_streaming_owned(output: impl Write + Send + 'static) -> Self {
Self::new(Box::new(SeekWrapper(output)))
}
}
impl<'a> BuilderExt<BorrowedOutput<'a>> {
pub fn new_seekable_borrowed(output: impl Write + Seek + 'a) -> Self {
Self::new(Box::new(output))
}
pub fn new_streaming_borrowed(output: impl Write + 'a) -> Self {
Self::new(Box::new(SeekWrapper(output)))
}
}
impl<W: Write + Seek> Clone for BuilderExt<W> {
fn clone(&self) -> Self {
Self {
tar: Arc::clone(&self.tar),
path: self.path.clone(),
}
}
}
impl<W: Write + Seek> BuilderExt<W> {
fn new(output: W) -> Self {
let enabled = Arc::new(AtomicBool::new(true));
Self {
tar: Arc::new(Mutex::new(BlowFuseOnDrop {
tar: Some(
tar::Builder::new(FusedWriteSeek {
output,
enabled: Arc::clone(&enabled),
})
.tap_mut(|tar| tar.sparse(true)),
),
enabled,
})),
path: PathBuf::new(),
}
}
pub fn descend(&self, subdir: &Path) -> io::Result<Self> {
Ok(Self {
tar: Arc::clone(&self.tar),
path: join_relative(&self.path, subdir)?,
})
}
pub fn blocking_write_fn<T>(
&self,
dst: &Path,
f: impl FnOnce(&mut tar::EntryWriter) -> T,
) -> io::Result<T> {
let dst = join_relative(&self.path, dst)?;
let mut header = tar::Header::new_gnu();
header.set_mode(0o644);
let mut tar = self.tar.blocking_lock();
let mut writer = tar.tar().append_writer(&mut header, dst)?;
let result = f(&mut writer);
writer.finish()?;
Ok(result)
}
pub fn blocking_append_file(&self, src: &Path, dst: &Path) -> io::Result<()> {
let dst = join_relative(&self.path, dst)?;
self.tar
.blocking_lock()
.tar()
.append_path_with_name(src, dst)
}
pub fn blocking_append_dir_all(&self, src: &Path, dst: &Path) -> io::Result<()> {
let dst = join_relative(&self.path, dst)?;
self.tar.blocking_lock().tar().append_dir_all(dst, src)
}
pub fn blocking_append_data(&self, src: &[u8], dst: &Path) -> io::Result<()> {
let dst = join_relative(&self.path, dst)?;
let mut header = tar::Header::new_gnu();
header.set_mode(0o644);
header.set_size(src.len() as u64);
self.tar
.blocking_lock()
.tar()
.append_data(&mut header, dst, src)
}
pub fn blocking_finish(self) -> io::Result<()> {
let mut bb: BlowFuseOnDrop<_> = Arc::try_unwrap(self.tar)
.map_err(|_| {
io::Error::other("finish called with multiple references to the tar builder")
})?
.into_inner();
let tar: tar::Builder<FusedWriteSeek<_>> = bb.tar.take().unwrap();
let mut wb: FusedWriteSeek<_> = tar.into_inner()?; wb.flush()?;
Ok(())
}
}
impl<W: Send + Write + Seek + 'static> BuilderExt<W> {
pub async fn append_data(&self, src: Vec<u8>, dst: &Path) -> io::Result<()> {
let dst = join_relative(&self.path, dst)?;
let mut header = tar::Header::new_gnu();
header.set_mode(0o644);
header.set_size(src.len() as u64);
self.run_async(move |tar| tar.append_data(&mut header, dst, src.as_slice()))
.await
}
pub async fn finish(self) -> io::Result<()> {
tokio::task::spawn_blocking(move || self.blocking_finish()).await?
}
async fn run_async<T, E>(
&self,
f: impl FnOnce(&mut tar::Builder<FusedWriteSeek<W>>) -> Result<T, E> + Send + 'static,
) -> Result<T, E>
where
T: Send + 'static,
E: Send + 'static + From<JoinError>,
{
let tar = Arc::clone(&self.tar);
tokio::task::spawn_blocking(move || f(tar.blocking_lock().tar())).await?
}
}
fn join_relative(base: &Path, rel_path: &Path) -> io::Result<PathBuf> {
if rel_path.is_absolute() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("path must be relative, but got {rel_path:?}"),
));
}
Ok(base.join(rel_path))
}
struct SeekWrapper<T>(T);
impl<T: Write> io::Write for SeekWrapper<T> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.0.write(buf)
}
fn flush(&mut self) -> io::Result<()> {
self.0.flush()
}
}
impl<T: Write> io::Seek for SeekWrapper<T> {
fn seek(&mut self, _: io::SeekFrom) -> io::Result<u64> {
Err(io::ErrorKind::NotSeekable.into())
}
}
#[cfg(test)]
mod tests {
use super::*;
struct DummyBridgeWriter(bool, Arc<Mutex<Vec<u8>>>);
impl Write for DummyBridgeWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
if self.0 {
return Err(io::Error::other("Forced error in write"));
}
self.1.blocking_lock().extend_from_slice(buf); Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
if self.0 {
return Err(io::Error::other("Forced error in flush"));
}
let _ = self.1.blocking_lock(); Ok(())
}
}
impl Seek for DummyBridgeWriter {
fn seek(&mut self, _: io::SeekFrom) -> io::Result<u64> {
unimplemented!()
}
}
#[tokio::test]
async fn test_dummy_finish_ok() {
let data = Arc::new(Mutex::new(Vec::new()));
let tar = BuilderExt::new_seekable_owned(DummyBridgeWriter(false, Arc::clone(&data)));
assert!(tar.finish().await.is_ok());
assert_eq!(data.lock().await.len(), 1024);
}
#[tokio::test]
async fn test_dummy_finish_fail() {
let data = Arc::new(Mutex::new(Vec::new()));
let tar = BuilderExt::new_seekable_owned(DummyBridgeWriter(true, Arc::clone(&data)));
assert!(tar.finish().await.is_err());
assert_eq!(data.lock().await.len(), 0);
}
#[tokio::test]
async fn test_dummy_drop_fail() {
let data = Arc::new(Mutex::new(Vec::new()));
let tar = BuilderExt::new_seekable_owned(DummyBridgeWriter(true, Arc::clone(&data)));
drop(tar);
assert_eq!(data.lock().await.len(), 0);
}
#[tokio::test]
async fn test_dummy_drop_ok() {
let data = Arc::new(Mutex::new(Vec::new()));
let tar = BuilderExt::new_seekable_owned(DummyBridgeWriter(false, Arc::clone(&data)));
drop(tar);
assert_eq!(data.lock().await.len(), 0);
}
#[test]
fn test_write_ok() {
let tar = BuilderExt::new_streaming_borrowed(Vec::new());
tar.blocking_append_data(b"foo", Path::new("foo")).unwrap();
tar.blocking_finish().unwrap();
}
#[test]
fn test_write_fail() {
let tar = BuilderExt::new_streaming_borrowed(Vec::new());
tar.blocking_append_data(b"foo", Path::new("foo")).unwrap();
let result = tar.blocking_write_fn(Path::new("foo"), |writer| writer.write_all(b"bar"));
assert_eq!(result.unwrap_err().kind(), io::ErrorKind::NotSeekable);
}
#[test]
fn test_writeseek_ok() {
let tar = BuilderExt::new_seekable_borrowed(io::Cursor::new(Vec::new()));
tar.blocking_append_data(b"foo", Path::new("foo")).unwrap();
tar.blocking_write_fn(Path::new("foo"), |writer| writer.write_all(b"bar"))
.unwrap()
.unwrap();
tar.blocking_finish().unwrap();
}
}