Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
149 changes: 140 additions & 9 deletions crates/walgit-store/src/s3.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,15 @@
//! ## Version tokens
//!
//! S3 `ETags` are used as opaque `Version` strings. Quotes are stripped
//! consistently on read and never stored. For non-multipart uploads the
//! `ETag` is the MD5 of the content; for multipart uploads it is a compound
//! hash. Callers never parse the token — equality comparison suffices.
//! consistently on read and never stored; conditional headers put them back
//! ([`entity_tag`]). For non-multipart uploads the `ETag` is the MD5 of the
//! content; for multipart uploads it is a compound hash. Callers never parse
//! the token — equality comparison suffices.
//!
//! ## Conditional PUT
//!
//! `PutMode::Create` → `If-None-Match: *` (object must not exist).
//! `PutMode::Update(v)` → `If-Match: <etag>` (CAS on current `ETag`).
//! `PutMode::Create` → `If-None-Match: *` (object must not exist).
//! `PutMode::Update(v)` → `If-Match: "<etag>"` (CAS on current `ETag`).
//! On failure the SDK returns a `PreconditionFailed` service error; we fill
//! `current` via a follow-up HEAD when the SDK doesn't include it.
//!
Expand Down Expand Up @@ -124,10 +125,10 @@ impl S3Store {
let mut builder = self.client.get_object().bucket(&self.bucket).key(key);

if let Some(v) = &opts.if_none_match {
builder = builder.if_none_match(v.as_str());
builder = builder.if_none_match(entity_tag(v));
}
if let Some(v) = &opts.if_match {
builder = builder.if_match(v.as_str());
builder = builder.if_match(entity_tag(v));
}
if let Some(r) = &opts.range {
builder = builder.range(Self::range_header(r));
Expand Down Expand Up @@ -393,7 +394,7 @@ impl ObjectStore for S3Store {
builder = builder.if_none_match("*");
}
PutMode::Update(v) => {
builder = builder.if_match(v.as_str());
builder = builder.if_match(entity_tag(v));
}
}

Expand Down Expand Up @@ -977,6 +978,20 @@ impl S3Store {
}
}

/// A version token as the entity-tag a conditional header carries. `ETag`s are
/// stored unquoted, but `If-Match` / `If-None-Match` take quoted entity-tags
/// (RFC 9110 §8.8.3, §13.1.1), the form every response's `ETag` header has.
/// AWS S3 and rustfs also accept a bare value; other S3-compatible stores do
/// not document it (Cloudflare R2 documents neither form).
fn entity_tag(v: &Version) -> String {
let s = v.as_str();
if s.len() >= 2 && s.starts_with('"') && s.ends_with('"') {
s.to_owned()
} else {
format!("\"{s}\"")
}
}

fn static_credentials(
access_key: &str,
secret_key: &str,
Expand All @@ -997,13 +1012,16 @@ fn static_credentials(
// 7. Multipart: CreateMultipartUpload + UploadPart + CompleteMultipartUpload
// supported. No conditional headers on Create/Complete (same as real S3).
// 8. ETags: quoted, MD5 for single-PUT, compound for multipart. Quotes
// stripped consistently in our Version.
// stripped consistently in our Version, restored on If-Match/If-None-Match
// (rustfs honours both forms).
// 9. force_path_style: required for rustfs local dev.

#[cfg(test)]
mod tests {
use super::*;

use std::sync::Arc;

use aws_sdk_s3::error::SdkError;
use aws_sdk_s3::operation::list_objects_v2::ListObjectsV2Error;
use aws_sdk_s3::operation::put_object::PutObjectError;
Expand Down Expand Up @@ -1141,6 +1159,119 @@ mod tests {
));
}

/// A fake S3 that records each request head and answers a PUT with 200 and
/// a GET with 304, both with a quoted `ETag` like real S3.
async fn recording_s3() -> (String, Arc<parking_lot::Mutex<Vec<String>>>) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let heads: Arc<parking_lot::Mutex<Vec<String>>> = Arc::default();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let endpoint = format!("http://{}", listener.local_addr().unwrap());
let recorded = heads.clone();
tokio::spawn(async move {
while let Ok((mut sock, _)) = listener.accept().await {
let recorded = recorded.clone();
tokio::spawn(async move {
let mut buf = Vec::new();
let mut chunk = [0u8; 8192];
let end = loop {
let n = sock.read(&mut chunk).await.unwrap_or(0);
if n == 0 {
return;
}
buf.extend_from_slice(&chunk[..n]);
if let Some(i) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
break i + 4;
}
};
let head = String::from_utf8_lossy(&buf[..end]).into_owned();
let body_len = head
.lines()
.filter_map(|l| l.split_once(':'))
.find(|(k, _)| k.eq_ignore_ascii_case("content-length"))
.and_then(|(_, v)| v.trim().parse::<usize>().ok())
.unwrap_or(0);
while buf.len() < end + body_len {
let n = sock.read(&mut chunk).await.unwrap_or(0);
if n == 0 {
break;
}
buf.extend_from_slice(&chunk[..n]);
}
let resp = if head.starts_with("GET ") {
"HTTP/1.1 304 Not Modified\r\nETag: \"abc\"\r\nConnection: close\r\n\r\n"
} else {
"HTTP/1.1 200 OK\r\nETag: \"def\"\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
};
recorded.lock().push(head);
let _ = sock.write_all(resp.as_bytes()).await;
});
}
});
(endpoint, heads)
}

fn header<'a>(head: &'a str, name: &str) -> Option<&'a str> {
head.lines()
.filter_map(|l| l.split_once(':'))
.find(|(k, _)| k.eq_ignore_ascii_case(name))
.map(|(_, v)| v.trim())
}

#[test]
fn entity_tags_are_quoted_once() {
assert_eq!(entity_tag(&Version::new("abc")), "\"abc\"");
assert_eq!(entity_tag(&Version::new("\"abc\"")), "\"abc\"");
assert_eq!(entity_tag(&Version::new("abc-3")), "\"abc-3\"");
}

/// Versions hold `ETag`s unquoted; on the wire a conditional header carries
/// the quoted entity-tag (RFC 9110), the form every `ETag` response has.
#[tokio::test]
async fn conditional_headers_carry_quoted_entity_tags() {
let (endpoint, heads) = recording_s3().await;
let store = S3Store {
client: client_for(&endpoint),
bucket: "b".into(),
http: reqwest::Client::new(),
multipart_threshold: u64::MAX,
multipart_part_size: 8 << 20,
};
let put = store
.put(
"manifest.pb",
PutBody::Bytes(Bytes::from_static(b"x")),
PutMode::Update(Version::new("abc")).into(),
)
.await
.unwrap();
assert_eq!(put.version.as_str(), "def", "stored unquoted");
let got = store
.get(
"manifest.pb",
GetOptions {
if_none_match: Some(Version::new("abc")),
..Default::default()
},
)
.await
.unwrap();
assert!(matches!(got, GetResult::NotModified { .. }));
let _ = store
.get(
"manifest.pb",
GetOptions {
if_match: Some(Version::new("abc")),
..Default::default()
},
)
.await;
let heads = heads.lock().clone();
assert_eq!(heads.len(), 3, "{heads:#?}");
assert_eq!(header(&heads[0], "if-match"), Some("\"abc\""));
assert_eq!(header(&heads[1], "if-none-match"), Some("\"abc\""));
assert_eq!(header(&heads[2], "if-match"), Some("\"abc\""));
}

#[test]
fn transient_codes_are_recognised() {
for code in [
Expand Down