perf(media-proxy): drop the extra HEAD on passthrough GETs (#2157)

This commit is contained in:
Hampus
2026-08-30 22:27:51 +02:00
committed by GitHub
parent cbf504dbb8
commit 1046edd903
3 changed files with 723 additions and 90 deletions
+119
View File
@@ -19,6 +19,13 @@ pub struct ContentRange {
pub size: Option<usize>,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum RequestRange<'a> {
Absent,
Forwardable(&'a str),
Unsatisfiable,
}
pub fn parse_range(header: Option<&str>, file_size: usize) -> ParsedRange {
let Some(raw) = header else {
return ParsedRange::default();
@@ -112,6 +119,52 @@ pub fn parse_bounded_request_range(header: Option<&str>, max_len: usize) -> Opti
Some(ByteRange { start, end })
}
pub fn classify_request_range(header: Option<&str>) -> RequestRange<'_> {
let Some(raw) = header else {
return RequestRange::Absent;
};
let trimmed = raw.trim_matches([' ', '\t']);
let Some(spec) = trimmed.strip_prefix("bytes=") else {
return RequestRange::Absent;
};
if spec.contains(',') {
return RequestRange::Absent;
}
let Some(dash) = spec.find('-') else {
return RequestRange::Absent;
};
let start_part = &spec[..dash];
let end_part = &spec[dash + 1..];
if start_part.is_empty() && end_part.is_empty() {
return RequestRange::Absent;
}
if start_part.is_empty() {
return match end_part.parse::<usize>() {
Ok(0) => RequestRange::Unsatisfiable,
Ok(_) => RequestRange::Forwardable(trimmed),
Err(_) => RequestRange::Absent,
};
}
let Ok(start) = start_part.parse::<usize>() else {
return RequestRange::Absent;
};
if end_part.is_empty() {
return RequestRange::Forwardable(trimmed);
}
match end_part.parse::<usize>() {
Ok(end) if end < start => RequestRange::Unsatisfiable,
Ok(_) => RequestRange::Forwardable(trimmed),
Err(_) => RequestRange::Absent,
}
}
pub fn parse_unsatisfiable_content_range(header: Option<&str>) -> Option<usize> {
let raw = header?;
let spec = raw.trim_matches([' ', '\t']).strip_prefix("bytes ")?;
let size_part = spec.strip_prefix('*')?.strip_prefix('/')?;
size_part.trim_matches([' ', '\t']).parse::<usize>().ok()
}
pub fn parse_content_range(header: Option<&str>) -> Option<ContentRange> {
let raw = header?;
let spec = raw.trim_matches([' ', '\t']).strip_prefix("bytes ")?;
@@ -237,6 +290,72 @@ mod tests {
);
}
#[test]
fn classified_request_ranges_agree_with_the_size_aware_parser() {
for raw in ["bytes=0-9", "bytes=-5", "bytes=50-", "bytes=0-9999"] {
assert_eq!(
RequestRange::Forwardable(raw),
classify_request_range(Some(raw)),
"range={raw} must reach the upstream"
);
assert!(parse_range(Some(raw), 100).range.is_some());
}
assert_eq!(
RequestRange::Forwardable("bytes=0-9"),
classify_request_range(Some(" bytes=0-9 "))
);
assert_eq!(
RequestRange::Forwardable("bytes=100-200"),
classify_request_range(Some("bytes=100-200")),
"only the object size can settle a range that starts past the end"
);
for raw in ["bytes=10-5", "bytes=-0"] {
assert_eq!(
RequestRange::Unsatisfiable,
classify_request_range(Some(raw)),
"range={raw} is unsatisfiable at every size"
);
assert!(parse_range(Some(raw), 100).unsatisfiable);
assert!(parse_range(Some(raw), 1).unsatisfiable);
}
for raw in [
"rows=0-9",
"bytes=",
"bytes=abc-def",
"bytes=0-abc",
"bytes=0-1, 2-3",
"bytes=0",
] {
assert_eq!(
RequestRange::Absent,
classify_request_range(Some(raw)),
"range={raw} must not reach the upstream"
);
assert_eq!(None, parse_range(Some(raw), 100).range);
assert!(!parse_range(Some(raw), 100).unsatisfiable);
}
assert_eq!(RequestRange::Absent, classify_request_range(None));
}
#[test]
fn unsatisfiable_content_range_parser_reads_the_total() {
assert_eq!(
Some(100),
parse_unsatisfiable_content_range(Some("bytes */100"))
);
assert_eq!(
Some(0),
parse_unsatisfiable_content_range(Some("bytes */0"))
);
assert_eq!(None, parse_unsatisfiable_content_range(Some("bytes */*")));
assert_eq!(
None,
parse_unsatisfiable_content_range(Some("bytes 0-9/100"))
);
assert_eq!(None, parse_unsatisfiable_content_range(Some("*/100")));
assert_eq!(None, parse_unsatisfiable_content_range(None));
}
#[test]
fn content_range_parser_accepts_known_and_unknown_totals() {
assert_eq!(
+382 -66
View File
@@ -12,7 +12,7 @@ use crate::{
request_log::{self, ErrorReason, Stage},
signing,
spool::{SpoolError, spool_to_temp},
storage::{HeadResult, RelayBody, RelayPutOptions, StorageError, Store, StreamObject},
storage::{RelayBody, RelayPutOptions, StorageError, Store, StreamObject},
timed_semaphore::TimedSemaphore,
upload_relay,
};
@@ -1769,6 +1769,56 @@ async fn serve_stored_passthrough_stream(
key: &str,
headers: &HeaderMap,
disposition: PassthroughDisposition<'_>,
) -> Response {
if method == Method::HEAD {
return serve_stored_passthrough_head(app, bucket, key, headers, &disposition).await;
}
if app.cfg.mode == DeploymentMode::Mp
&& image_extension_from_filename(key) == Some(AssetExtension::Svg)
{
return serve_stored_passthrough_svg(app, method, bucket, key, headers, &disposition).await;
}
let range_header = headers.get(header::RANGE).and_then(|v| v.to_str().ok());
let forwarded_range = match range::classify_request_range(range_header) {
range::RequestRange::Absent => None,
range::RequestRange::Forwardable(value) => Some(value),
range::RequestRange::Unsatisfiable => {
return passthrough_unsatisfiable_response(app, bucket, key, None).await;
}
};
let object = match app.store.stream_object(bucket, key, forwarded_range).await {
Ok(object) => object,
Err(err) => return storage_error_response(key, err),
};
if object.status == StatusCode::RANGE_NOT_SATISFIABLE {
return passthrough_unsatisfiable_response(app, bucket, key, object.total_length).await;
}
let content_type = passthrough_content_type(&object.content_type, key);
if app.cfg.mode == DeploymentMode::Mp && is_svg_content_type(&content_type) {
return serve_stored_passthrough_svg(app, method, bucket, key, headers, &disposition).await;
}
let total_len = match passthrough_total_len(app, bucket, key, object.total_length).await {
Ok(value) => value,
Err(err) => return storage_error_response(key, err),
};
if total_len > constants::MAX_MEDIA_PROXY_BYTES {
return storage_error_response(key, StorageError::StreamTooLong);
}
streaming_media_response(
method,
object,
total_len,
&content_type,
passthrough_disposition_header(&disposition, &content_type),
)
}
async fn serve_stored_passthrough_head(
app: &Arc<AppState>,
bucket: &str,
key: &str,
headers: &HeaderMap,
disposition: &PassthroughDisposition<'_>,
) -> Response {
let head = match app.store.head_object(bucket, key).await {
Ok(head) => head,
@@ -1777,25 +1827,13 @@ async fn serve_stored_passthrough_stream(
if head.content_length > constants::MAX_MEDIA_PROXY_BYTES as u64 {
return storage_error_response(key, StorageError::StreamTooLong);
}
let content_type = passthrough_content_type(&head, key);
let content_type = passthrough_content_type(&head.content_type, key);
if app.cfg.mode == DeploymentMode::Mp
&& (is_svg_content_type(&content_type)
|| image_extension_from_filename(key) == Some(AssetExtension::Svg))
{
let object = match app.store.read_object(bucket, key).await {
Ok(object) => object,
Err(err) => return storage_error_response(key, err),
};
let cache_identity = format!("{bucket}/{key}");
return serve_stored_svg_rasterized(
app,
method,
object.data,
&cache_identity,
headers,
&disposition,
)
.await;
return serve_stored_passthrough_svg(app, Method::HEAD, bucket, key, headers, disposition)
.await;
}
let total_len = match usize::try_from(head.content_length) {
Ok(value) => value,
@@ -1804,53 +1842,85 @@ async fn serve_stored_passthrough_stream(
let range_header = headers.get(header::RANGE).and_then(|v| v.to_str().ok());
let parsed_range = range::parse_range(range_header, total_len);
if parsed_range.unsatisfiable {
let mut response = Response::new(Body::empty());
*response.status_mut() = StatusCode::RANGE_NOT_SATISFIABLE;
http_headers::add_unsatisfiable_headers(response.headers_mut(), total_len);
return response;
return unsatisfiable_response(total_len);
}
let normalized_range = parsed_range
.range
.map(|r| format!("bytes={}-{}", r.start, r.end));
if method == Method::HEAD {
return passthrough_head_response(
&content_type,
total_len,
parsed_range.range,
passthrough_disposition_header(&disposition, &content_type),
);
}
let object = match app
.store
.stream_object(bucket, key, normalized_range.as_deref())
.await
{
Ok(object) => object,
Err(err) => return storage_error_response(key, err),
};
streaming_media_response(
method,
object,
passthrough_head_response(
&content_type,
total_len,
parsed_range.range,
&content_type,
passthrough_disposition_header(&disposition, &content_type),
passthrough_disposition_header(disposition, &content_type),
)
}
fn passthrough_content_type(head: &HeadResult, key: &str) -> String {
async fn serve_stored_passthrough_svg(
app: &Arc<AppState>,
method: Method,
bucket: &str,
key: &str,
headers: &HeaderMap,
disposition: &PassthroughDisposition<'_>,
) -> Response {
let object = match app.store.read_object(bucket, key).await {
Ok(object) => object,
Err(err) => return storage_error_response(key, err),
};
let cache_identity = format!("{bucket}/{key}");
serve_stored_svg_rasterized(
app,
method,
object.data,
&cache_identity,
headers,
disposition,
)
.await
}
async fn passthrough_total_len(
app: &Arc<AppState>,
bucket: &str,
key: &str,
known: Option<u64>,
) -> Result<usize, StorageError> {
let total = match known {
Some(value) => value,
None => app.store.head_object(bucket, key).await?.content_length,
};
usize::try_from(total).map_err(|_| StorageError::StreamTooLong)
}
async fn passthrough_unsatisfiable_response(
app: &Arc<AppState>,
bucket: &str,
key: &str,
known_total: Option<u64>,
) -> Response {
match passthrough_total_len(app, bucket, key, known_total).await {
Ok(total_len) => unsatisfiable_response(total_len),
Err(err) => storage_error_response(key, err),
}
}
fn unsatisfiable_response(total_len: usize) -> Response {
let mut response = Response::new(Body::empty());
*response.status_mut() = StatusCode::RANGE_NOT_SATISFIABLE;
http_headers::add_unsatisfiable_headers(response.headers_mut(), total_len);
response
}
fn passthrough_content_type(source_content_type: &str, key: &str) -> String {
let extension_mime = mime::extension_mime(key);
if extension_mime == Some("audio/mp4")
&& mime::normalize(Some(&head.content_type)) == Some("video/mp4")
&& mime::normalize(Some(source_content_type)) == Some("video/mp4")
{
return "audio/mp4".to_owned();
}
if content_type_is_trustworthy(&head.content_type) {
head.content_type.clone()
if content_type_is_trustworthy(source_content_type) {
source_content_type.to_owned()
} else {
extension_mime
.or_else(|| {
mime::normalize(Some(&head.content_type)).filter(|value| {
mime::normalize(Some(source_content_type)).filter(|value| {
!value.is_empty() && !value.eq_ignore_ascii_case("application/octet-stream")
})
})
@@ -1992,7 +2062,6 @@ fn streaming_media_response(
method: Method,
object: StreamObject,
total_len: usize,
byte_range: Option<range::ByteRange>,
content_type: &str,
disposition: Option<String>,
) -> Response {
@@ -2002,7 +2071,7 @@ fn streaming_media_response(
StatusCode::OK
};
let effective_byte_range = if status == StatusCode::PARTIAL_CONTENT {
byte_range
object.byte_range
} else {
None
};
@@ -3653,34 +3722,281 @@ mod tests {
assert_eq!("image/webp", mime::sniff(&raster.data).mime);
}
type PassthroughOriginRequests = Arc<tokio::sync::Mutex<Vec<(Method, Option<String>)>>>;
fn parse_origin_range(raw: &str) -> Option<(usize, usize)> {
let spec = raw.strip_prefix("bytes=")?;
let (start, end) = spec.split_once('-')?;
Some((start.parse().ok()?, end.parse().ok()?))
}
async fn passthrough_origin(
body: &'static [u8],
answer_416_with_total: bool,
) -> (String, PassthroughOriginRequests) {
use axum::body::Body;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let seen: PassthroughOriginRequests = Arc::new(tokio::sync::Mutex::new(Vec::new()));
let handler = Arc::clone(&seen);
let router = Router::new().fallback(any(move |request: axum::extract::Request| {
let seen = Arc::clone(&handler);
async move {
let (parts, _body) = request.into_parts();
let range = parts
.headers
.get(header::RANGE)
.and_then(|value| value.to_str().ok())
.map(ToOwned::to_owned);
seen.lock().await.push((parts.method, range.clone()));
let total = body.len();
let mut response = Response::new(Body::empty());
response
.headers_mut()
.insert(header::CONTENT_TYPE, HeaderValue::from_static("image/png"));
match range.as_deref().and_then(parse_origin_range) {
None => {
response
.headers_mut()
.insert(header::CONTENT_LENGTH, HeaderValue::from(total));
*response.body_mut() = Body::from(body);
}
Some((start, end)) if start < total => {
let end = end.min(total - 1);
*response.status_mut() = StatusCode::PARTIAL_CONTENT;
response.headers_mut().insert(
header::CONTENT_RANGE,
HeaderValue::from_str(&format!("bytes {start}-{end}/{total}")).unwrap(),
);
response
.headers_mut()
.insert(header::CONTENT_LENGTH, HeaderValue::from(end - start + 1));
*response.body_mut() = Body::from(&body[start..=end]);
}
Some(_) => {
*response.status_mut() = StatusCode::RANGE_NOT_SATISFIABLE;
if answer_416_with_total {
response.headers_mut().insert(
header::CONTENT_RANGE,
HeaderValue::from_str(&format!("bytes */{total}")).unwrap(),
);
}
}
}
response
}
}));
tokio::spawn(async move {
axum::serve(listener, router).await.unwrap();
});
(format!("http://{addr}"), seen)
}
fn s3_mode_test_config(endpoint: &str) -> Config {
Config::load_from_iter([
(
"FLUXER_MEDIA_PROXY_SECRET_KEY".to_owned(),
"secret".to_owned(),
),
("FLUXER_MEDIA_PROXY_MODE".to_owned(), "mp".to_owned()),
(
"FLUXER_MEDIA_PROXY_STORAGE_BACKEND".to_owned(),
"s3".to_owned(),
),
("FLUXER_S3_ENDPOINT".to_owned(), endpoint.to_owned()),
(
"FLUXER_S3_ACCESS_KEY_ID".to_owned(),
"AKIAIOSFODNN7EXAMPLE".to_owned(),
),
(
"FLUXER_S3_SECRET_ACCESS_KEY".to_owned(),
"wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY".to_owned(),
),
])
.unwrap()
}
async fn passthrough_request(cfg: Config, range: Option<&str>) -> Response {
use axum::{body::Body, http::Request};
let state = test_app_state(cfg);
let router = Router::new()
.fallback(any(catch_all))
.with_state(Arc::clone(&state));
let mut request = Request::builder().uri("/attachments/1/2/pic.png");
if let Some(range) = range {
request = request.header(header::RANGE, range);
}
tower::ServiceExt::oneshot(router, request.body(Body::empty()).unwrap())
.await
.unwrap()
}
#[tokio::test]
async fn passthrough_get_serves_the_whole_object_in_one_round_trip() {
let (endpoint, seen) = passthrough_origin(b"0123456789", true).await;
let response = passthrough_request(s3_mode_test_config(&endpoint), None).await;
assert_eq!(StatusCode::OK, response.status());
assert_eq!(
"10",
response.headers().get(header::CONTENT_LENGTH).unwrap()
);
assert_eq!(
"image/png",
response.headers().get(header::CONTENT_TYPE).unwrap()
);
assert!(response.headers().get(header::CONTENT_RANGE).is_none());
let body = axum::body::to_bytes(response.into_body(), 64)
.await
.unwrap();
assert_eq!(b"0123456789", &body[..]);
assert_eq!(vec![(Method::GET, None)], seen.lock().await.clone());
}
#[tokio::test]
async fn passthrough_ranged_get_reuses_the_upstream_content_range() {
let (endpoint, seen) = passthrough_origin(b"0123456789", true).await;
let response = passthrough_request(s3_mode_test_config(&endpoint), Some("bytes=2-5")).await;
assert_eq!(StatusCode::PARTIAL_CONTENT, response.status());
assert_eq!(
"bytes 2-5/10",
response.headers().get(header::CONTENT_RANGE).unwrap()
);
assert_eq!("4", response.headers().get(header::CONTENT_LENGTH).unwrap());
let body = axum::body::to_bytes(response.into_body(), 64)
.await
.unwrap();
assert_eq!(b"2345", &body[..]);
assert_eq!(
vec![(Method::GET, Some("bytes=2-5".to_owned()))],
seen.lock().await.clone()
);
}
#[tokio::test]
async fn passthrough_unsatisfiable_range_reuses_the_upstream_total() {
let (endpoint, seen) = passthrough_origin(b"0123456789", true).await;
let response =
passthrough_request(s3_mode_test_config(&endpoint), Some("bytes=20-30")).await;
assert_eq!(StatusCode::RANGE_NOT_SATISFIABLE, response.status());
assert_eq!(
"bytes */10",
response.headers().get(header::CONTENT_RANGE).unwrap()
);
assert_eq!(
vec![(Method::GET, Some("bytes=20-30".to_owned()))],
seen.lock().await.clone()
);
}
#[tokio::test]
async fn passthrough_falls_back_to_a_head_when_the_upstream_416_omits_the_total() {
let (endpoint, seen) = passthrough_origin(b"0123456789", false).await;
let response =
passthrough_request(s3_mode_test_config(&endpoint), Some("bytes=20-30")).await;
assert_eq!(StatusCode::RANGE_NOT_SATISFIABLE, response.status());
assert_eq!(
"bytes */10",
response.headers().get(header::CONTENT_RANGE).unwrap()
);
assert_eq!(
vec![
(Method::GET, Some("bytes=20-30".to_owned())),
(Method::HEAD, None)
],
seen.lock().await.clone()
);
}
#[tokio::test]
async fn passthrough_ignores_a_malformed_range_instead_of_forwarding_it() {
let (endpoint, seen) = passthrough_origin(b"0123456789", true).await;
let response =
passthrough_request(s3_mode_test_config(&endpoint), Some("bytes=abc-def")).await;
assert_eq!(StatusCode::OK, response.status());
assert_eq!(vec![(Method::GET, None)], seen.lock().await.clone());
}
#[tokio::test]
async fn passthrough_answers_a_reversed_range_without_asking_for_the_body() {
let (endpoint, seen) = passthrough_origin(b"0123456789", true).await;
let response =
passthrough_request(s3_mode_test_config(&endpoint), Some("bytes=10-5")).await;
assert_eq!(StatusCode::RANGE_NOT_SATISFIABLE, response.status());
assert_eq!(
"bytes */10",
response.headers().get(header::CONTENT_RANGE).unwrap()
);
assert_eq!(vec![(Method::HEAD, None)], seen.lock().await.clone());
}
#[tokio::test]
async fn passthrough_get_never_touches_the_write_endpoint() {
let (origin, origin_seen) = passthrough_origin(b"0123456789", true).await;
let (read, read_seen) = passthrough_origin(b"0123456789", true).await;
let mut cfg = s3_mode_test_config(&origin);
cfg.s3_read_endpoint = Some(read);
let response = passthrough_request(cfg, None).await;
assert_eq!(StatusCode::OK, response.status());
assert!(origin_seen.lock().await.is_empty());
assert_eq!(vec![(Method::GET, None)], read_seen.lock().await.clone());
}
#[tokio::test]
async fn passthrough_head_stays_a_single_round_trip() {
use axum::{body::Body, http::Request};
let (endpoint, seen) = passthrough_origin(b"0123456789", true).await;
let state = test_app_state(s3_mode_test_config(&endpoint));
let router = Router::new()
.fallback(any(catch_all))
.with_state(Arc::clone(&state));
let response = tower::ServiceExt::oneshot(
router,
Request::builder()
.method(Method::HEAD)
.uri("/attachments/1/2/pic.png")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(StatusCode::OK, response.status());
assert_eq!(
"10",
response.headers().get(header::CONTENT_LENGTH).unwrap()
);
assert_eq!(vec![(Method::HEAD, None)], seen.lock().await.clone());
}
#[test]
fn passthrough_content_type_preserves_non_media_metadata() {
let head = HeadResult {
content_length: 10,
content_type: "application/zip".to_owned(),
};
assert_eq!(
"application/zip",
passthrough_content_type(&head, "downloads/app.zip")
passthrough_content_type("application/zip", "downloads/app.zip")
);
}
#[test]
fn passthrough_content_type_prefers_known_extension_over_bad_metadata() {
let head = HeadResult {
content_length: 10,
content_type: "text/plain".to_owned(),
};
assert_eq!("image/png", passthrough_content_type(&head, "image.png"));
assert_eq!(
"image/png",
passthrough_content_type("text/plain", "image.png")
);
}
#[test]
fn passthrough_content_type_prefers_m4a_extension_over_mp4_metadata() {
let head = HeadResult {
content_length: 10,
content_type: "video/mp4".to_owned(),
};
assert_eq!("audio/mp4", passthrough_content_type(&head, "track.m4a"));
assert_eq!(
"audio/mp4",
passthrough_content_type("video/mp4", "track.m4a")
);
}
#[test]
+222 -24
View File
@@ -56,6 +56,8 @@ pub struct StreamObject {
pub status: StatusCode,
pub content_length: Option<u64>,
pub content_type: String,
pub byte_range: Option<crate::range::ByteRange>,
pub total_length: Option<u64>,
}
#[derive(Clone, Debug)]
@@ -334,6 +336,19 @@ impl Store {
}
let total_len = meta.len() as usize;
let parsed_range = crate::range::parse_range(range_header, total_len);
let content_type = mime::extension_mime(key)
.unwrap_or("application/octet-stream")
.to_owned();
if parsed_range.unsatisfiable {
return Ok(StreamObject {
content_length: Some(0),
body: Body::empty(),
status: StatusCode::RANGE_NOT_SATISFIABLE,
content_type,
byte_range: None,
total_length: Some(meta.len()),
});
}
let (status, body_len, start) = if let Some(r) = parsed_range.range {
(
StatusCode::PARTIAL_CONTENT,
@@ -352,9 +367,9 @@ impl Store {
content_length: Some(body_len),
body: Body::from_stream(ReaderStream::new(reader)),
status,
content_type: mime::extension_mime(key)
.unwrap_or("application/octet-stream")
.to_owned(),
content_type,
byte_range: parsed_range.range,
total_length: Some(meta.len()),
})
}
@@ -469,15 +484,10 @@ impl Store {
}
async fn head_s3(&self, bucket: &str, key: &str) -> Result<HeadResult, StorageError> {
let url = self.s3_url(bucket, key)?;
let signed = self.sign(Method::HEAD, &url, &[], None, &[])?;
let response = self
.client
.head(&url)
.headers(signed_headers(&signed, &self.cfg))
.send()
.await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
let url = self.s3_read_url(bucket, key)?;
let headers = self.read_headers(bucket, Method::HEAD, &url, &[])?;
let response = self.client.head(&url).headers(headers).send().await?;
if self.read_status_is_miss(bucket, response.status()) {
return Err(StorageError::NotFound);
}
if !response.status().is_success() {
@@ -526,10 +536,28 @@ impl Store {
if self.read_status_is_miss(bucket, response.status()) {
return Err(StorageError::NotFound);
}
if !response.status().is_success() {
let status = response.status();
let content_range = response
.headers()
.get(header::CONTENT_RANGE)
.and_then(|v| v.to_str().ok())
.map(ToOwned::to_owned);
if status == StatusCode::RANGE_NOT_SATISFIABLE {
return Ok(StreamObject {
body: Body::empty(),
status,
content_length: Some(0),
content_type: String::new(),
byte_range: None,
total_length: crate::range::parse_unsatisfiable_content_range(
content_range.as_deref(),
)
.map(|total| total as u64),
});
}
if !status.is_success() {
return Err(StorageError::S3(s3_error_summary(response).await));
}
let status = response.status();
let content_length: Option<u64> = response
.headers()
.get(header::CONTENT_LENGTH)
@@ -541,6 +569,22 @@ impl Store {
{
return Err(StorageError::StreamTooLong);
}
let (byte_range, total_length) = if status == StatusCode::PARTIAL_CONTENT {
let Some(parsed) = crate::range::parse_content_range(content_range.as_deref()) else {
return Err(StorageError::S3(
"partial response without a usable Content-Range".to_owned(),
));
};
(
Some(crate::range::ByteRange {
start: parsed.start,
end: parsed.end,
}),
parsed.size.map(|size| size as u64),
)
} else {
(None, content_length)
};
let content_type = response
.headers()
.get(header::CONTENT_TYPE)
@@ -552,6 +596,8 @@ impl Store {
status,
content_length,
content_type,
byte_range,
total_length,
})
}
@@ -1390,8 +1436,147 @@ mod tests {
assert_unsigned(&headers);
}
async fn range_server(
status: u16,
content_range: Option<&'static str>,
) -> (String, CapturedRequests) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let captured: CapturedRequests = std::sync::Arc::new(tokio::sync::Mutex::new(Vec::new()));
let handler = std::sync::Arc::clone(&captured);
let app = axum::Router::new().fallback(axum::routing::any(
move |request: axum::extract::Request| {
let captured = std::sync::Arc::clone(&handler);
async move {
let (parts, _body) = request.into_parts();
captured
.lock()
.await
.push((parts.method, parts.uri, parts.headers));
let mut response = axum::response::Response::new(Body::from("partial"));
*response.status_mut() = StatusCode::from_u16(status).unwrap();
if let Some(content_range) = content_range {
response.headers_mut().insert(
header::CONTENT_RANGE,
http::HeaderValue::from_static(content_range),
);
}
response
}
},
));
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
(format!("http://{addr}"), captured)
}
#[tokio::test]
async fn head_object_always_uses_origin_for_authoritative_length() {
async fn stream_s3_reads_the_span_and_the_total_off_a_partial_response() {
let (s3, _s3_seen) = range_server(206, Some("bytes 10-16/100")).await;
let tmp = tempfile::tempdir().unwrap();
let store = Store::new(s3_test_config(tmp.path(), &s3));
let object = store
.stream_object("cdn", "video.mp4", Some("bytes=10-16"))
.await
.unwrap();
assert_eq!(StatusCode::PARTIAL_CONTENT, object.status);
assert_eq!(
Some(crate::range::ByteRange { start: 10, end: 16 }),
object.byte_range
);
assert_eq!(Some(100), object.total_length);
}
#[tokio::test]
async fn stream_s3_rejects_a_partial_response_without_a_usable_content_range() {
let (s3, _s3_seen) = range_server(206, None).await;
let tmp = tempfile::tempdir().unwrap();
let store = Store::new(s3_test_config(tmp.path(), &s3));
let result = store
.stream_object("cdn", "video.mp4", Some("bytes=10-16"))
.await;
assert!(matches!(result, Err(StorageError::S3(_))));
}
#[tokio::test]
async fn stream_s3_surfaces_an_upstream_416_with_the_total_it_reports() {
let (s3, _s3_seen) = range_server(416, Some("bytes */100")).await;
let tmp = tempfile::tempdir().unwrap();
let store = Store::new(s3_test_config(tmp.path(), &s3));
let object = store
.stream_object("cdn", "video.mp4", Some("bytes=900-999"))
.await
.unwrap();
assert_eq!(StatusCode::RANGE_NOT_SATISFIABLE, object.status);
assert_eq!(None, object.byte_range);
assert_eq!(Some(100), object.total_length);
}
#[tokio::test]
async fn stream_s3_surfaces_an_upstream_416_that_omits_the_total() {
let (s3, _s3_seen) = range_server(416, None).await;
let tmp = tempfile::tempdir().unwrap();
let store = Store::new(s3_test_config(tmp.path(), &s3));
let object = store
.stream_object("cdn", "video.mp4", Some("bytes=900-999"))
.await
.unwrap();
assert_eq!(StatusCode::RANGE_NOT_SATISFIABLE, object.status);
assert_eq!(None, object.total_length);
}
#[tokio::test]
async fn local_stream_reports_an_unsatisfiable_range_instead_of_the_whole_file() {
let tmp = tempfile::tempdir().unwrap();
let store = Store::new(test_config(&tmp.path().canonicalize().unwrap()));
store
.write_object("cdn", "a/b.txt", b"hello world", "text/plain")
.await
.unwrap();
let object = store
.stream_object("cdn", "a/b.txt", Some("bytes=99-200"))
.await
.unwrap();
assert_eq!(StatusCode::RANGE_NOT_SATISFIABLE, object.status);
assert_eq!(Some(11), object.total_length);
assert_eq!(None, object.byte_range);
}
#[tokio::test]
async fn local_stream_reports_the_total_alongside_a_partial_span() {
let tmp = tempfile::tempdir().unwrap();
let store = Store::new(test_config(&tmp.path().canonicalize().unwrap()));
store
.write_object("cdn", "a/b.txt", b"hello world", "text/plain")
.await
.unwrap();
let object = store
.stream_object("cdn", "a/b.txt", Some("bytes=6-10"))
.await
.unwrap();
assert_eq!(StatusCode::PARTIAL_CONTENT, object.status);
assert_eq!(Some(11), object.total_length);
assert_eq!(
Some(crate::range::ByteRange { start: 6, end: 10 }),
object.byte_range
);
}
#[tokio::test]
async fn head_object_uses_the_read_endpoint_like_every_other_read() {
let (s3, s3_seen) = capture_server().await;
let (cdn, cdn_seen) = capture_server().await;
let tmp = tempfile::tempdir().unwrap();
@@ -1403,9 +1588,23 @@ mod tests {
store.head_object("cdn", "a.png").await.unwrap();
assert!(
cdn_seen.lock().await.is_empty(),
"HEAD must not hit the CDN"
s3_seen.lock().await.is_empty(),
"HEAD must not hit the origin"
);
let (method, uri, headers) = only_request(&cdn_seen).await;
assert_eq!(http::Method::HEAD, method);
assert_eq!("/a.png", uri.path());
assert_unsigned(&headers);
}
#[tokio::test]
async fn head_object_without_a_read_endpoint_still_signs_the_origin() {
let (s3, s3_seen) = capture_server().await;
let tmp = tempfile::tempdir().unwrap();
let store = Store::new(s3_test_config(tmp.path(), &s3));
store.head_object("cdn", "a.png").await.unwrap();
let (method, uri, headers) = only_request(&s3_seen).await;
assert_eq!(http::Method::HEAD, method);
assert_eq!("/cdn/a.png", uri.path());
@@ -1413,7 +1612,7 @@ mod tests {
}
#[tokio::test]
async fn body_reads_use_cdn_while_head_uses_origin() {
async fn body_reads_and_heads_both_use_the_cdn() {
let (s3, s3_seen) = capture_server().await;
let (cdn, cdn_seen) = capture_server().await;
let tmp = tempfile::tempdir().unwrap();
@@ -1428,13 +1627,12 @@ mod tests {
.await
.unwrap();
let s3_reqs = s3_seen.lock().await.clone();
assert!(s3_seen.lock().await.is_empty());
let cdn_reqs = cdn_seen.lock().await.clone();
assert_eq!(1, s3_reqs.len());
assert_eq!(http::Method::HEAD, s3_reqs[0].0);
assert_eq!(1, cdn_reqs.len());
assert_eq!(http::Method::GET, cdn_reqs[0].0);
assert_eq!("bytes=0-3", cdn_reqs[0].2.get(header::RANGE).unwrap());
assert_eq!(2, cdn_reqs.len());
assert_eq!(http::Method::HEAD, cdn_reqs[0].0);
assert_eq!(http::Method::GET, cdn_reqs[1].0);
assert_eq!("bytes=0-3", cdn_reqs[1].2.get(header::RANGE).unwrap());
}
#[tokio::test]