use crate::{ProgressEntry, ProgressListener, Pusher};
use bytes::Bytes;
use std::{
fs::File,
io::{Seek, Write},
};
use tokio::io::SeekFrom;
pub struct StdFilePusher {
file: File,
p: u64,
sync_all: bool,
listener: Option<ProgressListener>,
}
impl std::fmt::Debug for StdFilePusher {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StdFilePusher")
.field("p", &self.p)
.field("sync_all", &self.sync_all)
.finish_non_exhaustive()
}
}
impl StdFilePusher {
pub async fn new(file: tokio::fs::File, size: u64, sync_all: bool) -> std::io::Result<Self> {
file.set_len(size).await?;
Ok(Self {
file: file.into_std().await,
p: 0,
sync_all,
listener: None,
})
}
fn write_at(&mut self, start: u64, mut bytes: &[u8]) -> std::io::Result<()> {
if self.p != start {
if let Err(e) = self.file.seek(SeekFrom::Start(start)) {
self.p = u64::MAX;
return Err(e);
}
self.p = start;
}
while !bytes.is_empty() {
match self.file.write(bytes) {
Ok(0) => {
return Err(std::io::Error::new(
std::io::ErrorKind::WriteZero,
"failed to write any data",
));
}
Ok(n) => {
let old = self.p;
self.p += n as u64;
if let Some(l) = &mut self.listener {
l(old..self.p);
}
bytes = &bytes[n..];
}
Err(ref e) if e.kind() == std::io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
Ok(())
}
}
impl Pusher for StdFilePusher {
type Error = std::io::Error;
fn set_listener(&mut self, cb: ProgressListener) {
self.listener = Some(cb);
}
fn push(&mut self, range: &ProgressEntry, bytes: Bytes) -> Result<(), (Self::Error, Bytes)> {
if bytes.is_empty() {
return Ok(());
}
let start = range.start;
if let Err(e) = self.write_at(start, &bytes) {
#[allow(clippy::cast_possible_truncation)]
let written_len = if self.p >= start
&& let offset = (self.p - start) as usize
&& offset <= bytes.len()
{
offset
} else {
0
};
self.p = u64::MAX;
let remaining_bytes = if written_len < bytes.len() {
bytes.slice(written_len..)
} else {
Bytes::new()
};
return Err((e, remaining_bytes));
}
Ok(())
}
fn flush(&mut self) -> Result<(), Self::Error> {
if self.sync_all {
self.file.sync_all()?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
use super::*;
use crate::ProgressEntry;
use std::sync::{Arc, Mutex};
use std::{io::Read, vec::Vec};
use tempfile::NamedTempFile;
#[tokio::test]
async fn test_rand_file_pusher() {
let temp_file = NamedTempFile::new().unwrap();
let file_path = temp_file.path();
let mut pusher = StdFilePusher::new(temp_file.reopen().unwrap().into(), 10, false)
.await
.unwrap();
let data = b"234";
let range = 2..5;
pusher.push(&range, data[..].into()).unwrap();
pusher.flush().unwrap();
let mut file_content = Vec::new();
File::open(file_path)
.unwrap()
.read_to_end(&mut file_content)
.unwrap();
assert_eq!(file_content, b"\0\x00234\0\0\0\0\0");
}
#[tokio::test]
async fn test_debug_impl() {
let temp_file = NamedTempFile::new().unwrap();
let pusher = StdFilePusher::new(temp_file.reopen().unwrap().into(), 10, false)
.await
.unwrap();
let _ = format!("{pusher:?}");
}
#[tokio::test]
async fn test_noncontiguous_write_seeks_and_calls_listener() {
let temp_file = NamedTempFile::new().unwrap();
let mut pusher = StdFilePusher::new(temp_file.reopen().unwrap().into(), 10, false)
.await
.unwrap();
let seen = Arc::new(Mutex::new(None::<ProgressEntry>));
let seen2 = seen.clone();
pusher.set_listener(Box::new(move |r| {
*seen2.lock().unwrap() = Some(r);
}));
pusher.push(&(5..8), b"xyz"[..].into()).unwrap();
assert_eq!(*seen.lock().unwrap(), Some(5..8));
}
#[tokio::test]
async fn test_empty_push_is_noop() {
let temp_file = NamedTempFile::new().unwrap();
let mut pusher = StdFilePusher::new(temp_file.reopen().unwrap().into(), 10, false)
.await
.unwrap();
pusher.push(&(0..0), Bytes::new()).unwrap();
}
#[tokio::test]
async fn test_sync_all_flush() {
let temp_file = NamedTempFile::new().unwrap();
let mut pusher = StdFilePusher::new(temp_file.reopen().unwrap().into(), 10, true)
.await
.unwrap();
pusher.push(&(2..5), b"234"[..].into()).unwrap();
pusher.flush().unwrap();
}
}