use std::fs::File;
use std::io::{self, Read, Seek, SeekFrom};
use std::path::Path;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use std::time::{Duration, Instant};
const STALL_LIMIT: Duration = Duration::from_secs(30);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamStatus {
Downloading,
Complete,
Failed,
}
pub struct PartialFileSource {
file: File,
pos: u64,
bytes_written: Arc<crate::remote::downloads::ByteFeed>,
total: u64,
status: Arc<dyn Fn() -> StreamStatus + Send + Sync>,
stall_limit: Duration,
wait_for_bytes: bool,
advertise_len: bool,
whole_file_end: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProbeMode {
Full,
Lengthless,
LengthlessWholeEnd,
}
impl PartialFileSource {
pub fn open(
path: &Path,
bytes_written: Arc<crate::remote::downloads::ByteFeed>,
total: u64,
status: Arc<dyn Fn() -> StreamStatus + Send + Sync>,
mode: ProbeMode,
) -> io::Result<Self> {
let mut source = Self::with_stall_limit(path, bytes_written, total, status, STALL_LIMIT)?;
source.advertise_len = mode == ProbeMode::Full;
source.whole_file_end = mode != ProbeMode::Lengthless;
Ok(source)
}
pub fn open_for_probe(
path: &Path,
bytes_written: Arc<crate::remote::downloads::ByteFeed>,
total: u64,
status: Arc<dyn Fn() -> StreamStatus + Send + Sync>,
mode: ProbeMode,
) -> io::Result<Self> {
let mut source = Self::with_stall_limit(path, bytes_written, total, status, STALL_LIMIT)?;
source.wait_for_bytes = false;
source.advertise_len = mode == ProbeMode::Full;
source.whole_file_end = mode != ProbeMode::Lengthless;
Ok(source)
}
fn with_stall_limit(
path: &Path,
bytes_written: Arc<crate::remote::downloads::ByteFeed>,
total: u64,
status: Arc<dyn Fn() -> StreamStatus + Send + Sync>,
stall_limit: Duration,
) -> io::Result<Self> {
Ok(Self {
file: File::open(path)?,
pos: 0,
bytes_written,
total,
status,
stall_limit,
wait_for_bytes: true,
advertise_len: true,
whole_file_end: true,
})
}
fn available(&self) -> u64 {
let written = self.bytes_written.load(Ordering::Acquire);
match (self.status)() {
StreamStatus::Complete => self.file.metadata().map(|m| m.len()).unwrap_or(written),
_ => written,
}
}
fn read_available(&mut self, buf: &mut [u8], limit: u64) -> io::Result<usize> {
let to_read = (limit as usize).min(buf.len());
self.file.read(&mut buf[..to_read]).inspect(|n| {
self.pos += *n as u64;
})
}
}
impl Read for PartialFileSource {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let deadline = Instant::now() + self.stall_limit;
loop {
let available = self.available();
if available > self.pos {
let n = self.read_available(buf, available - self.pos)?;
if n > 0 {
return Ok(n);
}
}
match (self.status)() {
StreamStatus::Failed => {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"stream download failed before delivering the whole track",
));
}
StreamStatus::Complete if available <= self.pos => return Ok(0),
StreamStatus::Complete => {}
StreamStatus::Downloading => {
if self.total > 0 && available >= self.total && self.pos >= self.total {
return Ok(0);
}
}
}
if !self.wait_for_bytes {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"past what the download has delivered",
));
}
if Instant::now() >= deadline {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"stream download stalled",
));
}
self.bytes_written.wait_past(available, deadline);
}
}
}
impl Seek for PartialFileSource {
fn seek(&mut self, pos: SeekFrom) -> io::Result<u64> {
let target: i64 = match pos {
SeekFrom::Start(n) => n as i64,
SeekFrom::Current(n) => self.pos as i64 + n,
SeekFrom::End(n) => {
let len = if self.whole_file_end && self.total > 0 {
self.total
} else {
self.available()
};
len as i64 + n
}
};
if target < 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"seek before beginning of stream",
));
}
self.pos = self.file.seek(SeekFrom::Start(target as u64))?;
Ok(self.pos)
}
}
impl symphonia::core::io::MediaSource for PartialFileSource {
fn is_seekable(&self) -> bool {
true
}
fn byte_len(&self) -> Option<u64> {
(self.advertise_len && self.total > 0).then_some(self.total)
}
}
#[cfg(test)]
mod tests {
use std::io::Write;
use std::sync::atomic::AtomicU8;
use symphonia::core::io::MediaSource;
use super::*;
struct Fixture {
_dir: tempfile::TempDir,
path: std::path::PathBuf,
written: Arc<crate::remote::downloads::ByteFeed>,
status: Arc<AtomicU8>,
}
const DOWNLOADING: u8 = 0;
const COMPLETE: u8 = 1;
const FAILED: u8 = 2;
impl Fixture {
fn new() -> Self {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("track.opus.part");
File::create(&path).unwrap();
Self {
_dir: dir,
path,
written: crate::remote::downloads::ByteFeed::new(),
status: Arc::new(AtomicU8::new(DOWNLOADING)),
}
}
fn push(&self, chunk: &[u8]) {
let mut f = std::fs::OpenOptions::new()
.append(true)
.open(&self.path)
.unwrap();
f.write_all(chunk).unwrap();
f.flush().unwrap();
self.written.advance(chunk.len() as u64);
}
fn set(&self, status: u8) {
self.status.store(status, Ordering::Release);
self.written.done();
}
fn status_fn(&self) -> Arc<dyn Fn() -> StreamStatus + Send + Sync> {
let status = self.status.clone();
Arc::new(move || match status.load(Ordering::Acquire) {
COMPLETE => StreamStatus::Complete,
FAILED => StreamStatus::Failed,
_ => StreamStatus::Downloading,
})
}
fn source(&self, total: u64) -> PartialFileSource {
self.source_with_stall(total, STALL_LIMIT)
}
fn source_with_stall(&self, total: u64, stall: Duration) -> PartialFileSource {
let status = self.status.clone();
PartialFileSource::with_stall_limit(
&self.path,
self.written.clone(),
total,
Arc::new(move || match status.load(Ordering::Acquire) {
COMPLETE => StreamStatus::Complete,
FAILED => StreamStatus::Failed,
_ => StreamStatus::Downloading,
}),
stall,
)
.unwrap()
}
}
#[test]
fn reads_what_has_landed() {
let fx = Fixture::new();
fx.push(b"hello streaming world");
fx.set(COMPLETE);
let mut out = Vec::new();
fx.source(21).read_to_end(&mut out).unwrap();
assert_eq!(out, b"hello streaming world");
}
#[test]
fn read_stops_at_the_write_head_then_resumes() {
let fx = Fixture::new();
fx.push(b"abcd");
let mut src = fx.source(10);
let mut first = [0u8; 8];
assert_eq!(src.read(&mut first).unwrap(), 4);
assert_eq!(&first[..4], b"abcd");
std::thread::spawn({
let path = fx.path.clone();
let written = fx.written.clone();
move || {
std::thread::sleep(Duration::from_millis(20));
let mut f = std::fs::OpenOptions::new()
.append(true)
.open(&path)
.unwrap();
f.write_all(b"efghij").unwrap();
f.flush().unwrap();
written.advance(6);
}
});
let mut rest = [0u8; 8];
let n = src.read(&mut rest).unwrap();
assert_eq!(&rest[..n], b"efghij");
}
#[test]
fn seeks_freely_below_the_write_head() {
let fx = Fixture::new();
fx.push(b"0123456789");
let mut src = fx.source(1_000_000);
assert_eq!(src.seek(SeekFrom::Start(5)).unwrap(), 5);
let mut out = [0u8; 3];
src.read_exact(&mut out).unwrap();
assert_eq!(&out, b"567");
assert_eq!(src.seek(SeekFrom::Start(1)).unwrap(), 1);
src.read_exact(&mut out).unwrap();
assert_eq!(&out, b"123");
assert_eq!(src.seek(SeekFrom::Current(-2)).unwrap(), 2);
}
#[test]
fn seek_from_end_uses_the_advertised_length() {
let fx = Fixture::new();
fx.push(b"0123456789");
let mut src = fx.source(10);
assert_eq!(src.seek(SeekFrom::End(0)).unwrap(), 10);
assert_eq!(src.seek(SeekFrom::End(-3)).unwrap(), 7);
let mut out = [0u8; 3];
src.read_exact(&mut out).unwrap();
assert_eq!(&out, b"789");
}
#[test]
fn a_lengthless_source_ends_where_the_download_does() {
let fx = Fixture::new();
fx.push(b"0123456789");
let mut src = PartialFileSource::open(
&fx.path,
fx.written.clone(),
1_000,
fx.status_fn(),
ProbeMode::Lengthless,
)
.unwrap();
assert_eq!(src.byte_len(), None);
assert_eq!(src.seek(SeekFrom::End(0)).unwrap(), 10);
let mut src = PartialFileSource::open(
&fx.path,
fx.written.clone(),
1_000,
fx.status_fn(),
ProbeMode::Full,
)
.unwrap();
assert_eq!(src.seek(SeekFrom::End(0)).unwrap(), 1_000);
}
#[test]
fn an_ogg_source_still_ends_at_the_whole_file() {
let fx = Fixture::new();
fx.push(b"0123456789");
let mut src = PartialFileSource::open(
&fx.path,
fx.written.clone(),
1_000,
fx.status_fn(),
ProbeMode::LengthlessWholeEnd,
)
.unwrap();
assert_eq!(src.byte_len(), None);
assert_eq!(src.seek(SeekFrom::End(0)).unwrap(), 1_000);
}
#[test]
fn a_landed_download_ends_at_the_whole_file() {
let fx = Fixture::new();
fx.push(b"0123456789");
fx.set(COMPLETE);
let mut src = PartialFileSource::open(
&fx.path,
fx.written.clone(),
10,
fx.status_fn(),
ProbeMode::Lengthless,
)
.unwrap();
assert_eq!(src.seek(SeekFrom::End(0)).unwrap(), 10);
}
#[test]
fn seek_before_start_errors() {
let fx = Fixture::new();
fx.push(b"hello");
assert!(fx.source(5).seek(SeekFrom::Current(-1)).is_err());
}
#[test]
fn failed_download_errors_instead_of_reporting_eof() {
let fx = Fixture::new();
fx.push(b"partial");
fx.set(FAILED);
let mut src = fx.source(1000);
let mut out = [0u8; 7];
src.read_exact(&mut out).unwrap();
assert_eq!(&out, b"partial");
assert_eq!(
src.read(&mut out).unwrap_err().kind(),
io::ErrorKind::BrokenPipe
);
}
#[test]
fn failure_wakes_a_blocked_reader() {
let fx = Fixture::new();
let mut src = fx.source(1000);
let (status, written) = (fx.status.clone(), fx.written.clone());
std::thread::spawn(move || {
std::thread::sleep(Duration::from_millis(20));
status.store(FAILED, Ordering::Release);
written.done();
});
let mut out = [0u8; 8];
assert_eq!(
src.read(&mut out).unwrap_err().kind(),
io::ErrorKind::BrokenPipe
);
}
#[test]
fn a_probe_reads_only_what_has_arrived() {
let fx = Fixture::new();
fx.push(b"0123456789");
let mut src = PartialFileSource::open_for_probe(
&fx.path,
fx.written.clone(),
1_000,
fx.status_fn(),
ProbeMode::Full,
)
.unwrap();
let mut out = [0u8; 10];
src.read_exact(&mut out).unwrap();
assert_eq!(
src.read(&mut out).unwrap_err().kind(),
io::ErrorKind::UnexpectedEof,
"past the write head is an answer, not a wait"
);
}
#[test]
fn playback_reads_wait_for_what_has_not_arrived() {
let fx = Fixture::new();
fx.push(b"0123456789");
let mut src = PartialFileSource::open(
&fx.path,
fx.written.clone(),
1_000,
fx.status_fn(),
ProbeMode::Lengthless,
)
.unwrap();
let mut out = [0u8; 10];
src.read_exact(&mut out).unwrap();
std::thread::spawn({
let path = fx.path.clone();
let written = fx.written.clone();
move || {
std::thread::sleep(Duration::from_millis(20));
let mut f = std::fs::OpenOptions::new()
.append(true)
.open(&path)
.unwrap();
f.write_all(b"abcde").unwrap();
f.flush().unwrap();
written.advance(5);
}
});
let n = src.read(&mut out).unwrap();
assert_eq!(&out[..n], b"abcde", "it waited rather than giving up");
}
#[test]
fn stalled_download_times_out() {
let fx = Fixture::new();
let mut src = fx.source_with_stall(1000, Duration::from_millis(20));
let mut out = [0u8; 8];
assert_eq!(
src.read(&mut out).unwrap_err().kind(),
io::ErrorKind::TimedOut
);
}
#[test]
fn completion_ends_the_read_at_the_true_length() {
let fx = Fixture::new();
fx.push(b"chunked");
fx.set(COMPLETE);
let mut out = Vec::new();
fx.source(0).read_to_end(&mut out).unwrap();
assert_eq!(out, b"chunked");
}
#[test]
fn survives_the_part_file_being_renamed() {
let fx = Fixture::new();
fx.push(b"0123456789");
let mut src = fx.source(10);
let mut out = [0u8; 4];
src.read_exact(&mut out).unwrap();
assert_eq!(&out, b"0123");
std::fs::rename(&fx.path, fx.path.with_extension("")).unwrap();
fx.set(COMPLETE);
let mut rest = Vec::new();
src.read_to_end(&mut rest).unwrap();
assert_eq!(rest, b"456789");
}
#[test]
fn byte_len_is_the_advertised_length_only() {
let fx = Fixture::new();
assert_eq!(fx.source(42).byte_len(), Some(42));
assert_eq!(fx.source(0).byte_len(), None);
}
#[test]
fn is_seekable_true() {
let fx = Fixture::new();
assert!(fx.source(0).is_seekable());
}
}