use std::fs::File;
use std::io::{Read, Seek, SeekFrom};
use std::path::{Path, PathBuf};
use std::sync::Arc;
use crate::tiff_utils::AnyResult;
pub trait RangeReader: Send + Sync {
fn read_range(&self, offset: u64, length: usize) -> AnyResult<Vec<u8>>;
fn size(&self) -> u64;
fn identifier(&self) -> &str;
fn is_local(&self) -> bool {
let id = self.identifier();
!id.starts_with("http://") && !id.starts_with("https://") && !id.starts_with("s3://")
}
}
pub struct LocalRangeReader {
path: PathBuf,
size: u64,
}
impl LocalRangeReader {
pub fn new(path: impl AsRef<Path>) -> AnyResult<Self> {
let path = path.as_ref().to_path_buf();
let metadata = std::fs::metadata(&path)?;
Ok(Self {
path,
size: metadata.len(),
})
}
}
pub struct MemoryRangeReader {
data: Arc<Vec<u8>>,
identifier: String,
}
impl MemoryRangeReader {
#[must_use]
pub fn new(data: Vec<u8>, identifier: String) -> Self {
Self {
data: Arc::new(data),
identifier,
}
}
#[must_use]
pub fn from_arc(data: Arc<Vec<u8>>, identifier: String) -> Self {
Self { data, identifier }
}
}
impl RangeReader for MemoryRangeReader {
fn read_range(&self, offset: u64, length: usize) -> AnyResult<Vec<u8>> {
#[allow(clippy::cast_possible_truncation)]
let start = offset as usize;
let end = (start + length).min(self.data.len());
if start >= self.data.len() {
return Ok(vec![]);
}
Ok(self.data[start..end].to_vec())
}
fn size(&self) -> u64 {
self.data.len() as u64
}
fn identifier(&self) -> &str {
&self.identifier
}
fn is_local(&self) -> bool {
true }
}
impl RangeReader for LocalRangeReader {
fn read_range(&self, offset: u64, length: usize) -> AnyResult<Vec<u8>> {
let mut file = File::open(&self.path)?;
file.seek(SeekFrom::Start(offset))?;
let mut buffer = vec![0u8; length];
file.read_exact(&mut buffer)?;
Ok(buffer)
}
fn size(&self) -> u64 {
self.size
}
fn identifier(&self) -> &str {
self.path.to_str().unwrap_or("<invalid path>")
}
}
pub struct HttpRangeReader {
url: String,
size: u64,
client: reqwest::blocking::Client,
}
impl HttpRangeReader {
pub fn new(url: &str) -> AnyResult<Self> {
let client = reqwest::blocking::Client::builder()
.timeout(std::time::Duration::from_secs(30))
.build()?;
let response = client.head(url).send()?;
let size = response
.headers()
.get("content-length")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse().ok())
.ok_or_else(|| format!("HTTP server did not return Content-Length header for {url}"))?;
Ok(Self {
url: url.to_string(),
size,
client,
})
}
}
impl RangeReader for HttpRangeReader {
fn read_range(&self, offset: u64, length: usize) -> AnyResult<Vec<u8>> {
let range = format!("bytes={}-{}", offset, offset + length as u64 - 1);
let response = self.client
.get(&self.url)
.header("Range", range)
.send()?;
if !response.status().is_success() {
return Err(format!("HTTP request failed: {}", response.status()).into());
}
Ok(response.bytes()?.to_vec())
}
fn size(&self) -> u64 {
self.size
}
fn identifier(&self) -> &str {
&self.url
}
}
pub struct S3RangeReader {
size: u64,
url: String,
}
impl S3RangeReader {
pub fn new(url: &str) -> AnyResult<Self> {
let url_parsed = url::Url::parse(url)?;
if url_parsed.scheme() != "s3" {
return Err("URL must use s3:// scheme".into());
}
if url_parsed.host_str().is_none() {
return Err("Missing bucket in S3 URL".into());
}
let key = url_parsed.path().trim_start_matches('/');
if key.is_empty() {
return Err("Missing key in S3 URL".into());
}
Ok(Self {
size: 0,
url: url.to_string(),
})
}
pub fn from_https(url: &str) -> AnyResult<Self> {
let http_reader = HttpRangeReader::new(url)?;
Ok(Self {
size: http_reader.size,
url: url.to_string(),
})
}
}
impl RangeReader for S3RangeReader {
fn read_range(&self, offset: u64, length: usize) -> AnyResult<Vec<u8>> {
let client = reqwest::blocking::Client::new();
let range = format!("bytes={}-{}", offset, offset + length as u64 - 1);
let response = client
.get(&self.url)
.header("Range", range)
.send()?;
if !response.status().is_success() {
return Err(format!("S3 request failed: {}", response.status()).into());
}
Ok(response.bytes()?.to_vec())
}
fn size(&self) -> u64 {
self.size
}
fn identifier(&self) -> &str {
&self.url
}
}
pub fn create_range_reader(source: &str) -> AnyResult<Arc<dyn RangeReader>> {
if source.starts_with("s3://") {
Ok(Arc::new(crate::s3::S3RangeReaderSync::new(source)?))
} else if source.starts_with("http://") || source.starts_with("https://") {
Ok(Arc::new(HttpRangeReader::new(source)?))
} else {
Ok(Arc::new(LocalRangeReader::new(source)?))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
use tempfile::NamedTempFile;
#[test]
fn test_local_range_reader() {
let mut file = NamedTempFile::new().unwrap();
file.write_all(b"Hello, World!").unwrap();
let reader = LocalRangeReader::new(file.path()).unwrap();
assert_eq!(reader.size(), 13);
let data = reader.read_range(0, 5).unwrap();
assert_eq!(&data, b"Hello");
let data = reader.read_range(7, 5).unwrap();
assert_eq!(&data, b"World");
}
#[test]
fn test_memory_range_reader() {
let data = b"Hello, World!".to_vec();
let reader = MemoryRangeReader::new(data, "test://memory".to_string());
assert_eq!(reader.size(), 13);
assert_eq!(reader.identifier(), "test://memory");
assert!(reader.is_local());
let range1 = reader.read_range(0, 5).unwrap();
assert_eq!(&range1, b"Hello");
let range2 = reader.read_range(7, 5).unwrap();
assert_eq!(&range2, b"World");
let range3 = reader.read_range(10, 10).unwrap();
assert_eq!(&range3, b"ld!");
let range4 = reader.read_range(100, 10).unwrap();
assert!(range4.is_empty());
}
#[test]
fn test_memory_range_reader_from_arc() {
let data = Arc::new(b"Test data".to_vec());
let reader = MemoryRangeReader::from_arc(data.clone(), "arc://test".to_string());
assert_eq!(reader.size(), 9);
let result = reader.read_range(0, 4).unwrap();
assert_eq!(&result, b"Test");
}
}