use crate::config::ParseOptions;
use crate::error::{FileParseError, UnknownFormatError};
use crate::file::FileType;
use crate::probe::Probe;
use crate::util::math::F80;
use std::collections::VecDeque;
use std::fs::File;
use std::io::{BufReader, Cursor, ErrorKind, Read, Seek, SeekFrom, Write};
use std::ops::{Deref, DerefMut, Range};
use std::path::Path;
pub(crate) trait SeekStreamLen: Seek {
fn stream_len_hack(&mut self) -> std::io::Result<u64> {
use std::io::SeekFrom;
let current_pos = self.stream_position()?;
let len = self.seek(SeekFrom::End(0))?;
self.seek(SeekFrom::Start(current_pos))?;
Ok(len)
}
}
impl<T> SeekStreamLen for T where T: Seek {}
pub trait Truncate {
fn truncate(&mut self, new_len: u64) -> std::io::Result<()>;
}
impl Truncate for File {
fn truncate(&mut self, new_len: u64) -> std::io::Result<()> {
self.set_len(new_len)
}
}
impl Truncate for Vec<u8> {
fn truncate(&mut self, new_len: u64) -> std::io::Result<()> {
self.truncate(new_len as usize);
Ok(())
}
}
impl Truncate for VecDeque<u8> {
fn truncate(&mut self, new_len: u64) -> std::io::Result<()> {
self.truncate(new_len as usize);
Ok(())
}
}
impl<T> Truncate for Cursor<T>
where
T: Truncate,
{
fn truncate(&mut self, new_len: u64) -> std::io::Result<()> {
self.get_mut().truncate(new_len)
}
}
impl<T> Truncate for Box<T>
where
T: Truncate,
{
fn truncate(&mut self, new_len: u64) -> std::io::Result<()> {
self.as_mut().truncate(new_len)
}
}
impl<T> Truncate for &mut T
where
T: Truncate,
{
fn truncate(&mut self, new_len: u64) -> std::io::Result<()> {
(**self).truncate(new_len)
}
}
pub trait Length {
fn len(&self) -> std::io::Result<u64>;
}
impl Length for File {
fn len(&self) -> std::io::Result<u64> {
self.metadata().map(|m| m.len())
}
}
impl Length for Vec<u8> {
fn len(&self) -> std::io::Result<u64> {
Ok(self.len() as u64)
}
}
impl Length for VecDeque<u8> {
fn len(&self) -> std::io::Result<u64> {
Ok(self.len() as u64)
}
}
impl<T> Length for Cursor<T>
where
T: Length,
{
fn len(&self) -> std::io::Result<u64> {
Length::len(self.get_ref())
}
}
impl<T> Length for Box<T>
where
T: Length,
{
fn len(&self) -> std::io::Result<u64> {
Length::len(self.as_ref())
}
}
impl<T> Length for &T
where
T: Length,
{
fn len(&self) -> std::io::Result<u64> {
Length::len(*self)
}
}
impl<T> Length for &mut T
where
T: Length,
{
fn len(&self) -> std::io::Result<u64> {
Length::len(*self)
}
}
pub(crate) enum FileSource<'a, F> {
Ref(&'a mut F),
Path(BufReader<File>),
}
impl<F: FileLike> Length for FileSource<'_, F> {
fn len(&self) -> std::io::Result<u64> {
match self {
Self::Ref(file) => Length::len(file),
Self::Path(file) => Length::len(file.get_ref()),
}
}
}
impl<F: FileLike> Truncate for FileSource<'_, F> {
fn truncate(&mut self, size: u64) -> std::io::Result<()> {
match self {
Self::Ref(file) => Truncate::truncate(file, size),
Self::Path(file) => Truncate::truncate(file.get_mut(), size),
}
}
}
impl<F: FileLike> Read for FileSource<'_, F> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
match self {
Self::Ref(file) => Read::read(file, buf),
Self::Path(file) => Read::read(file, buf),
}
}
}
impl<F: FileLike> Write for FileSource<'_, F> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
match self {
Self::Ref(file) => Write::write(file, buf),
Self::Path(file) => Write::write(file.get_mut(), buf),
}
}
fn flush(&mut self) -> std::io::Result<()> {
match self {
Self::Ref(file) => Write::flush(file),
Self::Path(file) => Write::flush(file.get_mut()),
}
}
}
impl<F: FileLike> Seek for FileSource<'_, F> {
fn seek(&mut self, pos: SeekFrom) -> std::io::Result<u64> {
match self {
Self::Ref(file) => Seek::seek(file, pos),
Self::Path(file) => Seek::seek(file, pos),
}
}
}
pub(crate) struct VerifiedFile<'a, F> {
format: FileType,
file: FileSource<'a, F>,
}
impl<'a, F: FileLike> VerifiedFile<'a, F> {
pub(crate) fn new(
file: &'a mut F,
parse_options: ParseOptions,
) -> Result<Self, FileParseError> {
let probe = Probe::new(file).options(parse_options).guess_file_type()?;
match probe.file_type() {
Some(format) => Ok(Self {
format,
file: FileSource::Ref(probe.into_inner()),
}),
None => Err(UnknownFormatError.into()),
}
}
pub(crate) fn format(&self) -> FileType {
self.format
}
pub(crate) fn into_inner(self) -> FileSource<'a, F> {
self.file
}
}
impl VerifiedFile<'_, File> {
pub(crate) fn new_from_path(path: &Path) -> Result<Self, FileParseError> {
let probe = Probe::open(path)?;
match probe.file_type() {
Some(format) => Ok(Self {
format,
file: FileSource::Path(probe.into_inner()),
}),
None => Err(UnknownFormatError.into()),
}
}
}
impl<'a, F> Deref for VerifiedFile<'a, F> {
type Target = FileSource<'a, F>;
fn deref(&self) -> &Self::Target {
&self.file
}
}
impl<F> DerefMut for VerifiedFile<'_, F> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.file
}
}
pub trait FileLike: Read + Write + Seek + Truncate + Length {
fn splice(&mut self, range: Range<u64>, replacement: &impl AsRef<[u8]>) -> std::io::Result<()> {
replace_range(self, range, replacement.as_ref())
}
}
impl<T> FileLike for T where T: Read + Write + Seek + Truncate + Length {}
const MOVE_BUFFER_SIZE: usize = 64 * 1024;
fn replace_range<F>(file: &mut F, range: Range<u64>, replacement: &[u8]) -> std::io::Result<()>
where
F: FileLike + ?Sized,
{
if range.start > range.end {
return Err(std::io::Error::new(
ErrorKind::InvalidInput,
"range start exceeds range end",
));
}
let file_len = Length::len(file)?;
if range.end > file_len {
return Err(std::io::Error::new(
ErrorKind::InvalidInput,
"range extends beyond file length",
));
}
let old_len = range.end - range.start;
let replacement_len = u64::try_from(replacement.len())
.map_err(|_| std::io::Error::new(ErrorKind::InvalidInput, "replacement is too large"))?;
if range.end == file_len {
file.seek(SeekFrom::Start(range.start))?;
file.write_all(replacement)?;
if replacement_len < old_len {
file.truncate(old_len - replacement_len)?;
}
return Ok(());
}
let mut buffer = vec![0_u8; MOVE_BUFFER_SIZE];
match replacement_len.cmp(&old_len) {
std::cmp::Ordering::Greater => {
let difference = replacement_len - old_len;
extend_storage(file, difference, &buffer)?;
shift_right(file, range.end, file_len, difference, &mut buffer)?;
},
std::cmp::Ordering::Less => {
let difference = old_len - replacement_len;
shift_left(file, range.end, file_len, difference, &mut buffer)?;
file.truncate(file_len - difference)?;
},
std::cmp::Ordering::Equal => {},
}
file.seek(SeekFrom::Start(range.start))?;
file.write_all(replacement)?;
Ok(())
}
fn extend_storage<F>(file: &mut F, amount: u64, zeros: &[u8]) -> std::io::Result<()>
where
F: FileLike + ?Sized,
{
file.seek(SeekFrom::End(0))?;
let mut remaining = amount;
while remaining != 0 {
let chunk_len = usize::try_from(remaining.min(zeros.len() as u64))
.expect("chunk length is bounded by the in-memory buffer");
file.write_all(&zeros[..chunk_len])?;
remaining -= chunk_len as u64;
}
Ok(())
}
fn shift_right<F>(
file: &mut F,
start: u64,
end: u64,
amount: u64,
buffer: &mut [u8],
) -> std::io::Result<()>
where
F: FileLike + ?Sized,
{
let mut cursor = end;
while cursor > start {
let chunk_len = usize::try_from((cursor - start).min(buffer.len() as u64))
.expect("chunk length is bounded by the in-memory buffer");
let source = cursor - chunk_len as u64;
file.seek(SeekFrom::Start(source))?;
file.read_exact(&mut buffer[..chunk_len])?;
file.seek(SeekFrom::Start(source + amount))?;
file.write_all(&buffer[..chunk_len])?;
cursor = source;
}
Ok(())
}
fn shift_left<F>(
file: &mut F,
start: u64,
end: u64,
amount: u64,
buffer: &mut [u8],
) -> std::io::Result<()>
where
F: FileLike + ?Sized,
{
let mut cursor = start;
while cursor < end {
let chunk_len = usize::try_from((end - cursor).min(buffer.len() as u64))
.expect("chunk length is bounded by the in-memory buffer");
file.seek(SeekFrom::Start(cursor))?;
file.read_exact(&mut buffer[..chunk_len])?;
file.seek(SeekFrom::Start(cursor - amount))?;
file.write_all(&buffer[..chunk_len])?;
cursor += chunk_len as u64;
}
Ok(())
}
pub(crate) trait ReadExt: Read {
fn read_f80(&mut self) -> std::io::Result<F80>;
}
impl<R> ReadExt for R
where
R: Read,
{
fn read_f80(&mut self) -> std::io::Result<F80> {
let mut bytes = [0; 10];
self.read_exact(&mut bytes)?;
Ok(F80::from_be_bytes(bytes))
}
}
#[derive(Copy, Clone, Debug, Default, PartialEq)]
pub(crate) enum RevSearchStart {
#[default]
FromEnd,
FromCurrent,
}
#[derive(Copy, Clone, Debug, Default, PartialEq)]
pub(crate) enum RevSearchEnd {
StreamStart,
#[default]
FromCurrent,
Pos(u64),
}
pub(crate) struct RevPatternSearcher<'a, T> {
start: RevSearchStart,
end: RevSearchEnd,
buffer_size: u64,
pattern: &'a [u8],
reader: &'a mut T,
}
impl<T> RevPatternSearcher<'_, T>
where
T: Read + Seek,
{
pub(crate) fn buffer_size(&mut self, buffer_size: u64) -> &mut Self {
self.buffer_size = buffer_size;
self
}
pub(crate) fn start_pos(&mut self, start: RevSearchStart) -> &mut Self {
self.start = start;
self
}
pub(crate) fn end_pos(&mut self, end: RevSearchEnd) -> &mut Self {
self.end = end;
self
}
pub(crate) fn search(&mut self) -> std::io::Result<bool> {
if self.pattern.is_empty() {
return Ok(true);
}
let original_pos = self.reader.stream_position()?;
let pattern_len = self.pattern.len();
let start_pos = match self.start {
RevSearchStart::FromEnd => self.reader.seek(SeekFrom::End(0))?,
RevSearchStart::FromCurrent => original_pos,
};
let end_pos = match self.end {
RevSearchEnd::StreamStart => 0,
RevSearchEnd::FromCurrent => original_pos,
RevSearchEnd::Pos(p) => p,
};
if start_pos < end_pos
|| (start_pos - end_pos) < pattern_len as u64
|| self.buffer_size < pattern_len as u64
{
self.reader.seek(SeekFrom::Start(original_pos))?;
return Ok(false);
}
let overlap_step = self.buffer_size - ((pattern_len as u64) - 1);
let mut current_pos = start_pos;
let mut buf = vec![0; self.buffer_size as usize];
while current_pos > end_pos {
let window_size = current_pos - end_pos;
let read_size = std::cmp::min(self.buffer_size, window_size);
let read_start = current_pos - read_size;
self.reader.seek(SeekFrom::Start(read_start))?;
let window = &mut buf[..read_size as usize];
self.reader.read_exact(window)?;
if let Some(match_offset) = window
.windows(self.pattern.len())
.enumerate()
.rev()
.find_map(|(idx, window)| {
if window == self.pattern {
Some(idx)
} else {
None
}
}) {
self.reader
.seek(SeekFrom::Start(read_start + match_offset as u64))?;
return Ok(true);
}
current_pos -= std::cmp::min(read_size, overlap_step);
}
Ok(false)
}
}
pub(crate) trait ReadFindExt: Read + Seek + Sized {
fn rfind<'a>(&'a mut self, pattern: &'a [u8]) -> RevPatternSearcher<'a, Self> {
RevPatternSearcher {
start: RevSearchStart::default(),
end: RevSearchEnd::StreamStart,
buffer_size: 1024,
pattern,
reader: self,
}
}
}
impl<T> ReadFindExt for T where T: Read + Seek {}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{ParseOptions, WriteOptions};
use crate::file::AudioFile;
use crate::mpeg::MpegFile;
use crate::tag::Accessor;
use std::io::{Cursor, Read, Seek, SeekFrom, Write};
use std::iter::repeat_n;
use std::ops::Neg;
const TEST_ASSET: &str = "tests/files/assets/minimal/full_test.mp3";
fn test_asset_contents() -> Vec<u8> {
std::fs::read(TEST_ASSET).unwrap()
}
fn file() -> MpegFile {
let file_contents = test_asset_contents();
let mut reader = Cursor::new(file_contents);
MpegFile::read_from(&mut reader, ParseOptions::new()).unwrap()
}
fn alter_tag(file: &mut MpegFile) {
let tag = file.id3v2_mut().unwrap();
tag.set_artist(String::from("Bar artist"));
}
fn revert_tag(file: &mut MpegFile) {
let tag = file.id3v2_mut().unwrap();
tag.set_artist(String::from("Foo artist"));
}
#[test_log::test]
fn io_save_to_file() {
let mut file = file();
alter_tag(&mut file);
let mut temp_file = tempfile::tempfile().unwrap();
let file_content = std::fs::read(TEST_ASSET).unwrap();
temp_file.write_all(&file_content).unwrap();
temp_file.rewind().unwrap();
file.save_to(&mut temp_file, WriteOptions::new().preferred_padding(0))
.expect("Failed to save to file");
temp_file.rewind().unwrap();
let mut file = MpegFile::read_from(&mut temp_file, ParseOptions::new()).unwrap();
revert_tag(&mut file);
temp_file.rewind().unwrap();
file.save_to(&mut temp_file, WriteOptions::new().preferred_padding(0))
.expect("Failed to save to file");
temp_file.rewind().unwrap();
let mut current_file_contents = Vec::new();
temp_file.read_to_end(&mut current_file_contents).unwrap();
assert_eq!(current_file_contents, test_asset_contents());
}
#[test_log::test]
fn io_save_to_vec() {
let mut file = file();
alter_tag(&mut file);
let file_content = std::fs::read(TEST_ASSET).unwrap();
let mut reader = Cursor::new(file_content);
file.save_to(&mut reader, WriteOptions::new().preferred_padding(0))
.expect("Failed to save to vec");
reader.rewind().unwrap();
let mut file = MpegFile::read_from(&mut reader, ParseOptions::new()).unwrap();
revert_tag(&mut file);
reader.rewind().unwrap();
file.save_to(&mut reader, WriteOptions::new().preferred_padding(0))
.expect("Failed to save to vec");
let current_file_contents = reader.into_inner();
assert_eq!(current_file_contents, test_asset_contents());
}
#[test_log::test]
fn io_save_using_references() {
struct File {
buf: Vec<u8>,
}
let mut f = File {
buf: std::fs::read(TEST_ASSET).unwrap(),
};
let mut file = file();
alter_tag(&mut file);
{
let mut reader = Cursor::new(&mut f.buf);
file.save_to(&mut reader, WriteOptions::new().preferred_padding(0))
.expect("Failed to save to vec");
}
{
let mut reader = Cursor::new(&f.buf[..]);
file = MpegFile::read_from(&mut reader, ParseOptions::new()).unwrap();
revert_tag(&mut file);
}
{
let mut reader = Cursor::new(&mut f.buf);
file.save_to(&mut reader, WriteOptions::new().preferred_padding(0))
.expect("Failed to save to vec");
}
let current_file_contents = f.buf;
assert_eq!(current_file_contents, test_asset_contents());
}
#[test_log::test]
fn rev_search() {
const PAT: &[u8] = b"PATTERN";
let mut data1 = PAT.to_vec();
data1.extend(repeat_n(0, 5000));
let mut stream1 = Cursor::new(data1);
assert!(stream1.rfind(PAT).search().unwrap());
let mut data2 = PAT.to_vec();
data2.extend(repeat_n(0, 1023));
let mut stream2 = Cursor::new(data2);
assert!(stream2.rfind(PAT).search().unwrap());
let mut data3 = PAT.to_vec();
let junk_len = 20;
data3.extend(repeat_n(0, junk_len));
data3.extend(PAT);
data3.extend(repeat_n(0, junk_len));
let last_occurence_offset = data3.len() - (junk_len + PAT.len());
let mut stream3 = Cursor::new(data3);
assert!(stream3.rfind(PAT).search().unwrap());
assert_eq!(stream3.position(), last_occurence_offset as u64);
let mut data4 = PAT.to_vec();
data4.extend(repeat_n(0, junk_len));
data4.extend(PAT);
data4.extend(repeat_n(0, junk_len));
data4.extend(PAT);
let middle_match_offset = PAT.len() + junk_len;
let mut stream4 = Cursor::new(data4);
stream4
.seek(SeekFrom::End(((PAT.len() - 3) as i64).neg()))
.unwrap();
assert!(
stream4
.rfind(PAT)
.start_pos(RevSearchStart::FromCurrent)
.end_pos(RevSearchEnd::StreamStart)
.search()
.unwrap()
);
assert_eq!(stream4.position(), middle_match_offset as u64);
}
fn apply_range(input: Vec<u8>, range: Range<usize>, replacement: &[u8]) {
let mut expected = input.clone();
drop(expected.splice(range.clone(), replacement.iter().copied()));
let mut cursor = Cursor::new(input);
replace_range(
&mut cursor,
(range.start as u64)..(range.end as u64),
replacement,
)
.expect("range replacement should succeed");
let actual = cursor.into_inner();
assert_eq!(actual, expected);
}
#[test]
fn replace_range_equal_size() {
apply_range(b"0123456789".to_vec(), 2..5, b"XYZ");
}
#[test]
fn replace_range_grows() {
apply_range(b"0123456789".to_vec(), 2..5, b"abcdef");
}
#[test]
fn replace_range_shrinks() {
apply_range(b"0123456789".to_vec(), 2..8, b"X");
}
#[test]
fn replace_range_grows_across_multiple_buffers() {
let mut input = b"prefix".to_vec();
input.extend((0..(MOVE_BUFFER_SIZE * 3 + 17)).map(|index| (index % 251) as u8));
apply_range(input, 1..4, b"a much longer metadata replacement");
}
#[test]
fn replace_range_shrinks_across_multiple_buffers() {
let mut input = b"prefix".to_vec();
input.extend((0..(MOVE_BUFFER_SIZE * 3 + 17)).map(|index| (index % 251) as u8));
apply_range(input, 1..4, b"x");
}
}