use reqwest::{
blocking::{Client, Response},
header::HeaderValue,
};
use std::io::{self, Read};
const MAX_RETRIES: u32 = 5;
#[cfg(not(test))]
const BASE_RETRY_DELAY_MS: u64 = 200;
#[cfg(test)]
const BASE_RETRY_DELAY_MS: u64 = 1;
fn parse_content_range_start(value: &str) -> Option<u64> {
let mut parts = value.split_whitespace();
let unit = parts.next()?;
if !unit.eq_ignore_ascii_case("bytes") {
return None;
}
let range = parts.next()?;
let start = range.split_once('-')?.0;
start.trim().parse().ok()
}
enum RangeNotSatisfiable {
Complete,
Incomplete,
}
fn classify_range_not_satisfiable(offset: u64, content_length: Option<u64>) -> RangeNotSatisfiable {
match content_length {
Some(len) if offset >= len => RangeNotSatisfiable::Complete,
_ => RangeNotSatisfiable::Incomplete,
}
}
enum ValidatorCheck {
Match,
Modified,
Unverifiable,
}
fn compare_validators(
original_last_modified: Option<&HeaderValue>,
original_etag: Option<&HeaderValue>,
resume_last_modified: Option<&HeaderValue>,
resume_etag: Option<&HeaderValue>,
) -> ValidatorCheck {
if original_last_modified.is_none() && original_etag.is_none() {
return ValidatorCheck::Match;
}
let mut shared = false;
if let (Some(original), Some(resume)) = (original_etag, resume_etag) {
shared = true;
if original != resume {
return ValidatorCheck::Modified;
}
}
if let (Some(original), Some(resume)) = (original_last_modified, resume_last_modified) {
shared = true;
if original != resume {
return ValidatorCheck::Modified;
}
}
if shared {
ValidatorCheck::Match
} else {
ValidatorCheck::Unverifiable
}
}
pub(crate) struct ResumableHttpReader {
client: Client,
url: String,
response: Response,
offset: u64,
content_length: Option<u64>,
last_modified: Option<HeaderValue>,
etag: Option<HeaderValue>,
}
enum Resume {
Resumed,
Eof,
Unsupported,
Failed,
}
impl ResumableHttpReader {
pub fn new(client: Client, url: String, response: Response) -> Self {
let content_length: Option<u64> = response
.headers()
.get(reqwest::header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse().ok());
let last_modified = response
.headers()
.get(reqwest::header::LAST_MODIFIED)
.cloned();
let etag = response.headers().get(reqwest::header::ETAG).cloned();
Self {
client,
url,
response,
offset: 0,
content_length,
last_modified,
etag,
}
}
fn resume(&mut self) -> io::Result<Resume> {
for attempt in 0..MAX_RETRIES {
let resp = match self
.client
.get(&self.url)
.header(reqwest::header::RANGE, format!("bytes={}-", self.offset))
.header(reqwest::header::ACCEPT_ENCODING, "identity")
.send()
{
Ok(resp) => resp,
Err(_) => {
let backoff_ms = BASE_RETRY_DELAY_MS.saturating_mul(1u64 << attempt.min(4));
std::thread::sleep(std::time::Duration::from_millis(backoff_ms));
continue;
}
};
return match resp.status() {
reqwest::StatusCode::OK if self.offset == 0 => {
self.response = resp;
Ok(Resume::Resumed)
}
reqwest::StatusCode::RANGE_NOT_SATISFIABLE => {
match classify_range_not_satisfiable(self.offset, self.content_length) {
RangeNotSatisfiable::Complete => Ok(Resume::Eof),
RangeNotSatisfiable::Incomplete => Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
format!(
"server returned 416 Range Not Satisfiable while resuming at byte {}, \
but the transfer is incomplete (declared content length: {}); \
refusing to return truncated data",
self.offset,
self.content_length
.map(|len| len.to_string())
.unwrap_or_else(|| "unknown".to_string()),
),
)),
}
}
reqwest::StatusCode::PARTIAL_CONTENT => {
self.accept_resumed_response(resp)?;
Ok(Resume::Resumed)
}
_ => Ok(Resume::Unsupported),
};
}
Ok(Resume::Failed)
}
fn accept_resumed_response(&mut self, resp: Response) -> io::Result<()> {
let content_range = resp
.headers()
.get(reqwest::header::CONTENT_RANGE)
.ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
"resumed response is missing the Content-Range header",
)
})?
.to_str()
.map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
let start = parse_content_range_start(content_range).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("malformed Content-Range header: {content_range}"),
)
})?;
if start != self.offset {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("server resumed at byte {start}, expected {}", self.offset),
));
}
let resume_last_modified = resp.headers().get(reqwest::header::LAST_MODIFIED);
let resume_etag = resp.headers().get(reqwest::header::ETAG);
match compare_validators(
self.last_modified.as_ref(),
self.etag.as_ref(),
resume_last_modified,
resume_etag,
) {
ValidatorCheck::Match => {}
ValidatorCheck::Modified => {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"resumed resource validators (ETag/Last-Modified) do not match the original; \
the resource changed mid-transfer",
));
}
ValidatorCheck::Unverifiable => {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"resumed response omitted the ETag/Last-Modified validators the original \
provided; cannot confirm the resource is unchanged",
));
}
}
self.response = resp;
Ok(())
}
}
impl Read for ResumableHttpReader {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let mut stalled_retries = 0u32;
loop {
match self.response.read(buf) {
Ok(0) => return Ok(0),
Ok(n) => {
self.offset += n as u64;
return Ok(n);
}
Err(original_err) => {
if stalled_retries >= MAX_RETRIES {
return Err(original_err);
}
stalled_retries += 1;
match self.resume()? {
Resume::Resumed => continue,
Resume::Eof => return Ok(0),
Resume::Unsupported | Resume::Failed => return Err(original_err),
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc,
};
use std::{
io::prelude::*,
net::{TcpListener, TcpStream},
thread,
};
use crate::resumable_http::{ResumableHttpReader, MAX_RETRIES};
fn read_request(stream: &mut TcpStream) -> String {
let mut data = Vec::new();
let mut byte = [0u8; 1];
while let Ok(1) = stream.read(&mut byte) {
data.push(byte[0]);
if data.ends_with(b"\r\n\r\n") {
break;
}
}
String::from_utf8_lossy(&data).to_string()
}
#[test]
fn no_drop() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
read_request(&mut stream);
let response = "HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n1234567890";
stream.write_all(response.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
reader.read_to_string(&mut buf).unwrap();
assert_eq!(buf.as_str(), "1234567890");
handle.join().unwrap();
}
#[test]
fn drop_resume() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 = "HTTP/1.1 200 OK\r\nContent-Length: 10\r\nLast-Modified: Tue, 15 Nov 1994 12:45:26 GMT\r\nETag: \"v1\"\r\n\r\n12345";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let (mut stream2, _) = listener.accept().unwrap();
let req = read_request(&mut stream2);
assert!(req.to_ascii_lowercase().contains("range: bytes=5-"));
let response_part2 = "HTTP/1.1 206 Partial Content\r\nContent-Length: 5\r\nContent-Range: bytes 5-9/10\r\nLast-Modified: Tue, 15 Nov 1994 12:45:26 GMT\r\nETag: \"v1\"\r\n\r\n67890";
stream2.write_all(response_part2.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
reader.read_to_string(&mut buf).unwrap();
assert_eq!(buf.as_str(), "1234567890");
handle.join().unwrap();
}
#[test]
fn resume_200_null_offset() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 =
"HTTP/1.1 200 OK\r\nContent-Length: 10\r\nConnection: close\r\n\r\n";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let (mut stream2, _) = listener.accept().unwrap();
let req = read_request(&mut stream2);
assert!(req.to_ascii_lowercase().contains("range: bytes=0-"));
let response_part2 = "HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n1234567890";
stream2.write_all(response_part2.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
reader.read_to_string(&mut buf).unwrap();
assert_eq!(buf.as_str(), "1234567890");
handle.join().unwrap();
}
#[test]
fn range_not_supported_is_err() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 = "HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n12345";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let (mut stream2, _) = listener.accept().unwrap();
read_request(&mut stream2);
let response_part2 = "HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n1234567890";
stream2.write_all(response_part2.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
assert!(reader.read_to_string(&mut buf).is_err());
handle.join().unwrap();
}
#[test]
fn range_oob_is_err() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 = "HTTP/1.1 200 OK\r\nContent-Length: 20\r\n\r\n1234567890";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let (mut stream2, _) = listener.accept().unwrap();
let req = read_request(&mut stream2);
assert!(req.to_ascii_lowercase().contains("range: bytes=10-"));
let response_part2 = "HTTP/1.1 416 Range Not Satisfiable\r\nContent-Length: 0\r\n\r\n";
stream2.write_all(response_part2.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
assert!(reader.read_to_string(&mut buf).is_err());
handle.join().unwrap();
}
#[test]
fn classify_416_only_eof_when_complete() {
use crate::resumable_http::{classify_range_not_satisfiable, RangeNotSatisfiable};
assert!(matches!(
classify_range_not_satisfiable(10, Some(10)),
RangeNotSatisfiable::Complete
));
assert!(matches!(
classify_range_not_satisfiable(11, Some(10)),
RangeNotSatisfiable::Complete
));
assert!(matches!(
classify_range_not_satisfiable(5, Some(10)),
RangeNotSatisfiable::Incomplete
));
assert!(matches!(
classify_range_not_satisfiable(5, None),
RangeNotSatisfiable::Incomplete
));
}
#[test]
fn compare_validators_decision_table() {
use crate::resumable_http::{compare_validators, ValidatorCheck};
use reqwest::header::HeaderValue;
let etag_a = HeaderValue::from_static("\"v1\"");
let etag_b = HeaderValue::from_static("\"v2\"");
let lm_a = HeaderValue::from_static("Tue, 15 Nov 1994 12:45:26 GMT");
let lm_b = HeaderValue::from_static("Tue, 15 Nov 1995 12:45:26 GMT");
assert!(matches!(
compare_validators(None, None, Some(&lm_b), Some(&etag_b)),
ValidatorCheck::Match
));
assert!(matches!(
compare_validators(Some(&lm_a), None, Some(&lm_a), None),
ValidatorCheck::Match
));
assert!(matches!(
compare_validators(None, Some(&etag_a), None, Some(&etag_a)),
ValidatorCheck::Match
));
assert!(matches!(
compare_validators(Some(&lm_a), None, Some(&lm_b), None),
ValidatorCheck::Modified
));
assert!(matches!(
compare_validators(None, Some(&etag_a), None, Some(&etag_b)),
ValidatorCheck::Modified
));
assert!(matches!(
compare_validators(None, Some(&etag_a), Some(&lm_a), Some(&etag_a)),
ValidatorCheck::Match
));
assert!(matches!(
compare_validators(Some(&lm_a), None, Some(&lm_a), Some(&etag_a)),
ValidatorCheck::Match
));
assert!(matches!(
compare_validators(Some(&lm_a), Some(&etag_a), None, Some(&etag_a)),
ValidatorCheck::Match
));
assert!(matches!(
compare_validators(Some(&lm_a), Some(&etag_a), Some(&lm_b), Some(&etag_a)),
ValidatorCheck::Modified
));
assert!(matches!(
compare_validators(Some(&lm_a), Some(&etag_a), None, None),
ValidatorCheck::Unverifiable
));
assert!(matches!(
compare_validators(Some(&lm_a), None, None, Some(&etag_a)),
ValidatorCheck::Unverifiable
));
}
#[test]
fn new_last_modified_is_err() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 =
"HTTP/1.1 200 OK\r\nContent-Length: 10\r\nLast-Modified: Tue, 15 Nov 1994 12:45:26 GMT\r\n\r\n12345";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let (mut stream2, _) = listener.accept().unwrap();
let req = read_request(&mut stream2);
assert!(req.to_ascii_lowercase().contains("range: bytes=5-"));
let response_part2 = "HTTP/1.1 206 Partial Content\r\nContent-Length: 5\r\nContent-Range: bytes 5-9/10\r\nLast-Modified: Tue, 15 Nov 1995 12:45:26 GMT\r\n\r\n67890";
stream2.write_all(response_part2.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
assert!(reader.read_to_string(&mut buf).is_err());
handle.join().unwrap();
}
#[test]
fn etag_mismatch_is_err() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 =
"HTTP/1.1 200 OK\r\nContent-Length: 10\r\nETag: \"v1\"\r\n\r\n12345";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let (mut stream2, _) = listener.accept().unwrap();
let req = read_request(&mut stream2);
assert!(req.to_ascii_lowercase().contains("range: bytes=5-"));
let response_part2 = "HTTP/1.1 206 Partial Content\r\nContent-Length: 5\r\nContent-Range: bytes 5-9/10\r\nETag: \"v2\"\r\n\r\n67890";
stream2.write_all(response_part2.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
assert!(reader.read_to_string(&mut buf).is_err());
handle.join().unwrap();
}
#[test]
fn missing_last_modified_on_resume_is_err() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 =
"HTTP/1.1 200 OK\r\nContent-Length: 10\r\nLast-Modified: Tue, 15 Nov 1994 12:45:26 GMT\r\n\r\n12345";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let (mut stream2, _) = listener.accept().unwrap();
let req = read_request(&mut stream2);
assert!(req.to_ascii_lowercase().contains("range: bytes=5-"));
let response_part2 = "HTTP/1.1 206 Partial Content\r\nContent-Length: 5\r\nContent-Range: bytes 5-9/10\r\n\r\n67890";
stream2.write_all(response_part2.as_bytes()).unwrap();
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
assert!(reader.read_to_string(&mut buf).is_err());
handle.join().unwrap();
}
#[test]
fn max_retries_exhausted_is_err() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let attempts_counter = Arc::new(AtomicUsize::new(0));
let server_counter = attempts_counter.clone();
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 = "HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n12345";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
loop {
let (mut stream, _) = listener.accept().unwrap();
let req = read_request(&mut stream);
if req.is_empty() || req.starts_with("STOP") {
break;
}
server_counter.fetch_add(1, Ordering::SeqCst);
drop(stream);
}
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
assert!(reader.read_to_string(&mut buf).is_err());
if let Ok(mut wake_stream) = TcpStream::connect(format!("127.0.0.1:{}", port)) {
let _ = wake_stream.write_all(b"STOP");
}
handle.join().unwrap();
assert_eq!(
attempts_counter.load(Ordering::SeqCst),
MAX_RETRIES as usize
);
}
#[test]
fn no_data_206() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/data.txt", port);
let attempts_counter = Arc::new(AtomicUsize::new(0));
let server_counter = attempts_counter.clone();
let handle = thread::spawn(move || {
let (mut stream1, _) = listener.accept().unwrap();
read_request(&mut stream1);
let response_part1 = "HTTP/1.1 200 OK\r\nContent-Length: 10\r\n\r\n12345";
stream1.write_all(response_part1.as_bytes()).unwrap();
drop(stream1);
let response_206 = "HTTP/1.1 206 Partial Content\r\nContent-Length: 5\r\nContent-Range: bytes 5-9/10\r\n\r\n";
loop {
let (mut stream, _) = listener.accept().unwrap();
let req = read_request(&mut stream);
if req.is_empty() || req.starts_with("STOP") {
break;
}
server_counter.fetch_add(1, Ordering::SeqCst);
stream.write_all(response_206.as_bytes()).unwrap();
drop(stream);
}
});
let client = reqwest::blocking::Client::new();
let resp = client.get(&url).send().unwrap();
let mut reader = ResumableHttpReader::new(client, url, resp);
let mut buf = String::new();
assert!(reader.read_to_string(&mut buf).is_err());
if let Ok(mut wake_stream) = TcpStream::connect(format!("127.0.0.1:{}", port)) {
let _ = wake_stream.write_all(b"STOP");
}
handle.join().unwrap();
assert_eq!(
attempts_counter.load(Ordering::SeqCst),
MAX_RETRIES as usize
);
}
}