diff --git a/server/src/handlers.rs b/server/src/handlers.rs index 4eee1761d0..003649b949 100644 --- a/server/src/handlers.rs +++ b/server/src/handlers.rs @@ -5,14 +5,14 @@ use actix_web::{ body::SizedStream, dev::ServiceRequest, http::{ - self, - header::{self, ContentLength, ContentType, EntityTag, HttpDate}, + self, StatusCode, + header::{self, ContentLength, ContentType, EntityTag, HttpDate, Range}, }, web::{Data, Header, Path, Payload}, }; use aws_sdk_s3::error::SdkError; use chrono::{DateTime, Utc}; -use futures::{StreamExt, stream}; +use futures::StreamExt; use lockable::LockPool; use serde::{Deserialize, Serialize}; use size::Size; @@ -181,6 +181,13 @@ async fn extract_headers(request: &mut ServiceRequest) -> HandlerResult<(Headers )) } +async fn extract_range_header(request: &mut ServiceRequest) -> Option { + request + .extract::>() + .await + .map(|header| header.0.to_string()) + .ok() +} #[derive(Serialize, Deserialize, Debug)] pub struct PartData { workspace: Uuid, @@ -397,6 +404,8 @@ pub async fn get(request: HttpRequest) -> HandlerResult { let etag = objectpart_etag(&parts).unwrap(); let date = objectpart_date(&parts).unwrap(); + let range = extract_range_header(&mut request).await; + match none_match(request.request(), Some(etag.clone()))? { Some(false) => HttpResponse::NotModified() .insert_header((header::ETAG, etag)) @@ -419,10 +428,31 @@ pub async fn get(request: HttpRequest) -> HandlerResult { } } + let accept_ranges = objectpart_accept_ranges(&parts); + if let Some(accept_ranges) = accept_ranges { + response.insert_header((header::ACCEPT_RANGES, accept_ranges)); + } + response.insert_header((header::ETAG, etag)); response.insert_header((header::LAST_MODIFIED, HttpDate::from(date))); response.insert_header((header::CACHE_CONTROL, CONFIG.cache_control.clone())); - response.body(merge::stream(s3, parts).await?) + + match range { + Some(range) => { + let partial = merge::partial(s3, parts, range).await?; + + if partial.partial { + response.status(StatusCode::PARTIAL_CONTENT); + } + + if let Some(content_range) = partial.content_range { + response.insert_header((header::CONTENT_RANGE, content_range)); + } + + response.body(partial.stream) + } + None => response.body(merge::stream(s3, parts).await?), + } } } } else { @@ -467,6 +497,11 @@ pub async fn head(request: HttpRequest) -> HandlerResult { } } + let accept_ranges = objectpart_accept_ranges(&parts); + if let Some(accept_ranges) = accept_ranges { + response.insert_header((header::ACCEPT_RANGES, accept_ranges)); + } + response.insert_header((header::ETAG, etag)); response.insert_header((header::LAST_MODIFIED, HttpDate::from(date))); response.insert_header((header::CACHE_CONTROL, CONFIG.cache_control.clone())); @@ -507,6 +542,20 @@ fn objectpart_strategy(parts: &Vec>) -> Option>) -> Option<&str> { + let strategy = objectpart_strategy(parts)?; + match strategy { + MergeStrategy::JsonPatch => None, + _ => { + if parts.len() == 1 { + Some("bytes") + } else { + None + } + } + } +} + fn validate_patch_conditionals( req: &HttpRequest, parts: &Vec>, diff --git a/server/src/merge.rs b/server/src/merge.rs index 827dafe55e..4eee6496d5 100644 --- a/server/src/merge.rs +++ b/server/src/merge.rs @@ -85,6 +85,44 @@ pub fn validate_patch_body(merge_strategy: MergeStrategy, blob: &Blob) -> Handle } } +pub struct PartialResponse { + pub partial: bool, + pub content_range: Option, + pub stream: SizedStream>>>>, +} + +#[instrument(level = "debug", skip_all)] +pub async fn partial( + s3: Arc, + parts: Vec>, + range: String, +) -> anyhow::Result { + let part = parts.first().unwrap(); + + let mut response = s3 + .get_object() + .bucket(&CONFIG.s3_bucket) + .key(&part.data.blob) + .range(range) + .send() + .await?; + + let content_range = response.content_range().map(|s| s.to_string()); + let content_length = response.content_length().map_or(0, |c| c as u64); + + let stream = stream! { + while let Some(chunk) = response.body.next().await { + yield Ok(Bytes::from(chunk?)); + }; + }; + + Ok(PartialResponse { + partial: part.data.size != content_length as usize, + content_range, + stream: SizedStream::new(content_length, Box::pin(stream)), + }) +} + #[instrument(level = "debug", skip_all)] pub async fn stream( s3: Arc, diff --git a/tests/src/get.rs b/tests/src/get.rs index 22e5de0f54..dde207b920 100644 --- a/tests/src/get.rs +++ b/tests/src/get.rs @@ -84,3 +84,73 @@ pub async fn get_conditional() -> eyre::Result<()> { Ok(()) } + +#[tanu::test] +pub async fn get_partial() -> eyre::Result<()> { + let key = random_key(); + let text = random_text(1024 * 1024 * 5); + + let http = Client::new(); + + let res = http.key_put(&key).body(text.clone()).send().await?; + check!(res.status().is_success()); + + let res = http.key_get(&key).send().await?; + check!(res.status().is_success()); + check_eq!(res.header("accept-ranges"), Some("bytes")); + + let res = http + .key_get(&key) + .header("range", "bytes=0-127") + .send() + .await?; + check_eq!(res.status(), http::StatusCode::PARTIAL_CONTENT); + check_eq!(res.header("content-length"), Some("128")); + check_eq!(res.header("content-range"), Some("bytes 0-127/5242880")); + + let res = http + .key_get(&key) + .header("range", "bytes=0-5242879") + .send() + .await?; + check_eq!(res.status(), http::StatusCode::OK); + check_eq!(res.header("content-length"), Some("5242880")); + check_eq!(res.header("content-range"), Some("bytes 0-5242879/5242880")); + + Ok(()) +} + +#[tanu::test] +pub async fn get_partial_inline() -> eyre::Result<()> { + let key = random_key(); + let text = random_text(1024); + + let http = Client::new(); + + let res = http.key_put(&key).body(text.clone()).send().await?; + check!(res.status().is_success()); + + let res = http.key_get(&key).send().await?; + check!(res.status().is_success()); + check_eq!(res.header("accept-ranges"), Some("bytes")); + + let res = http + .key_get(&key) + .header("range", "bytes=0-31") + .send() + .await?; + check_eq!(res.status(), http::StatusCode::PARTIAL_CONTENT); + check_eq!(res.header("content-length"), Some("32")); + check_eq!(res.header("content-range"), Some("bytes 0-31/1024")); + + let res = http + .key_get(&key) + .header("range", "bytes=0-1023") + .send() + .await?; + check_eq!(res.status(), http::StatusCode::OK); + check_eq!(res.header("content-length"), Some("1024")); + check_eq!(res.header("content-range"), Some("bytes 0-1023/1024")); + + Ok(()) +}