From bd571a575f38aa4da5cb3b792d78f9c5c3ea4726 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Mon, 1 Jun 2026 07:40:20 +0800 Subject: [PATCH] fix(io-core,signer): replace unwrap() with proper error handling (#3150) * fix(io-core,signer): replace unwrap() with proper error handling ## io-core (issue #653 item 1) - pool.rs: replace 8 .lock().unwrap() with poisoned recovery - pool.rs: replace semaphore acquire unwrap with graceful fallback - deadlock_detector.rs: replace 5 .lock().unwrap() with match/ok ## signer (issue #653 item 1) - Add SignV2Error enum, try_pre_sign_v2, try_sign_v2 - Replace 14 unwrap() in v2 signing with ? propagation - Add try_streaming_sign_v4, replace 5 expect("err") with descriptive errors - get_host_addr returns Result instead of panicking All backward-compat wrappers preserved. 95 io-core + 18 signer tests pass. * fix: address signer and pool review comments * test: update signer v2 string-to-sign test * fix: address signer and pool review followups * fix: tighten signer host and pool fallback --- crates/io-core/src/deadlock_detector.rs | 16 +- crates/io-core/src/pool.rs | 57 +++-- crates/signer/src/lib.rs | 4 + .../signer/src/request_signature_streaming.rs | 174 +++++++++++++-- crates/signer/src/request_signature_v2.rs | 207 +++++++++++++++--- crates/signer/src/utils.rs | 36 ++- 6 files changed, 421 insertions(+), 73 deletions(-) diff --git a/crates/io-core/src/deadlock_detector.rs b/crates/io-core/src/deadlock_detector.rs index facf328c8..778220e09 100644 --- a/crates/io-core/src/deadlock_detector.rs +++ b/crates/io-core/src/deadlock_detector.rs @@ -155,7 +155,7 @@ impl DeadlockDetector { /// Register a new lock. pub fn register_lock(&self, lock_type: LockType) -> u64 { let id = { - let mut next = self.next_lock_id.lock().unwrap(); + let mut next = self.next_lock_id.lock().unwrap_or_else(|e| e.into_inner()); *next += 1; *next }; @@ -243,7 +243,10 @@ impl DeadlockDetector { return None; } - let graph = self.wait_graph.lock().unwrap(); + let graph = match self.wait_graph.lock() { + Ok(g) => g, + Err(_) => return None, + }; // Build adjacency list let mut adj: HashMap> = HashMap::new(); @@ -310,7 +313,10 @@ impl DeadlockDetector { return Vec::new(); } - let locks = self.locks.lock().unwrap(); + let locks = match self.locks.lock() { + Ok(l) => l, + Err(_) => return Vec::new(), + }; let mut result = Vec::new(); for (&id, info) in locks.iter() { @@ -349,13 +355,13 @@ impl DeadlockDetector { /// Get lock info. pub fn get_lock_info(&self, lock_id: u64) -> Option { - let locks = self.locks.lock().unwrap(); + let locks = self.locks.lock().ok()?; locks.get(&lock_id).cloned() } /// Get total number of registered locks. pub fn lock_count(&self) -> usize { - let locks = self.locks.lock().unwrap(); + let locks = self.locks.lock().unwrap_or_else(|e| e.into_inner()); locks.len() } } diff --git a/crates/io-core/src/pool.rs b/crates/io-core/src/pool.rs index c6d4567dc..a8d7577fb 100644 --- a/crates/io-core/src/pool.rs +++ b/crates/io-core/src/pool.rs @@ -226,8 +226,9 @@ impl BytesPool { pub async fn acquire_buffer(&self, size: usize) -> PooledBuffer { let tier = self.select_tier(size); let mut buffer = tier.acquire_buffer(size, &self.metrics).await; - // Set tier reference for return on drop - buffer.tier = Some(Arc::clone(tier)); + if buffer._permit.is_some() { + buffer.tier = Some(Arc::clone(tier)); + } buffer } @@ -304,12 +305,12 @@ impl PoolTier { } fn set_metrics(&self, metrics: Arc) { - *self.metrics.lock().unwrap() = Some(metrics); + *self.metrics.lock().unwrap_or_else(|e| e.into_inner()) = Some(metrics); } fn take_or_allocate_buffer(&self, size: usize, pool_metrics: &BytesPoolMetrics) -> (BytesMut, bool) { let buffer_opt = { - let mut available = self.available_buffers.lock().unwrap(); + let mut available = self.available_buffers.lock().unwrap_or_else(|e| e.into_inner()); available.pop() }; let was_reused = buffer_opt.is_some(); @@ -368,10 +369,20 @@ impl PoolTier { async fn acquire_buffer(&self, size: usize, pool_metrics: &BytesPoolMetrics) -> PooledBuffer { // Acquire semaphore permit (owned for storage in PooledBuffer) - let permit = Arc::clone(&self.semaphore).acquire_owned().await.unwrap(); + let permit = match Arc::clone(&self.semaphore).acquire_owned().await { + Ok(p) => p, + Err(_) => { + let buffer = BytesMut::with_capacity(size.max(self.buffer_size)); + return PooledBuffer { + buffer: ManuallyDrop::new(buffer), + tier: None, + _permit: None, + }; + } + }; // Use the pool's shared metrics for recording - let _metrics_lock = self.metrics.lock().unwrap(); + let _metrics_lock = self.metrics.lock().unwrap_or_else(|e| e.into_inner()); let _metrics = _metrics_lock.as_ref().unwrap(); // Record acquisition @@ -394,7 +405,7 @@ impl PoolTier { let permit = Arc::clone(&self.semaphore).try_acquire_owned().ok()?; // Use the pool's shared metrics for recording - let _metrics_lock = self.metrics.lock().unwrap(); + let _metrics_lock = self.metrics.lock().unwrap_or_else(|e| e.into_inner()); let _metrics = _metrics_lock.as_ref().unwrap(); // Record acquisition @@ -414,11 +425,11 @@ impl PoolTier { /// Return a buffer to the pool for reuse. fn return_buffer(&self, buffer: BytesMut) { - let mut available = self.available_buffers.lock().unwrap(); + let mut available = self.available_buffers.lock().unwrap_or_else(|e| e.into_inner()); // Limit the size of the pool to prevent unbounded growth if available.len() < self.max_buffers { available.push(buffer); - if let Some(ref metrics) = *self.metrics.lock().unwrap() { + if let Some(ref metrics) = *self.metrics.lock().unwrap_or_else(|e| e.into_inner()) { metrics.available_buffers.fetch_add(1, Ordering::Relaxed); } } else { @@ -428,7 +439,7 @@ impl PoolTier { Some(current.saturating_sub(released_bytes)) }) .ok(); - if let Some(ref metrics) = *self.metrics.lock().unwrap() { + if let Some(ref metrics) = *self.metrics.lock().unwrap_or_else(|e| e.into_inner()) { metrics .current_allocated_bytes .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| { @@ -447,11 +458,10 @@ impl Drop for PooledBuffer { // buffer moves it exactly once into the pool when a tier still owns it. #[allow(unsafe_code)] fn drop(&mut self) { - // Return buffer to pool if tier reference exists + // Return buffer to pool if tier reference exists. + // Otherwise, drop the standalone fallback buffer normally. + let buffer = unsafe { ManuallyDrop::take(&mut self.buffer) }; if let Some(ref tier) = self.tier { - // SAFETY: We're in drop(), so this is the last use of the buffer - // ManuallyDrop allows us to take the value without running BytesMut's drop - let buffer = unsafe { ManuallyDrop::take(&mut self.buffer) }; tier.return_buffer(buffer); } // The permit is automatically dropped here, releasing the semaphore slot @@ -503,7 +513,10 @@ impl std::fmt::Debug for PoolTier { .field("buffer_size", &self.buffer_size) .field("max_buffers", &self.max_buffers) .field("available_permits", &self.semaphore.available_permits()) - .field("available_buffers", &self.available_buffers.lock().unwrap().len()) + .field( + "available_buffers", + &self.available_buffers.lock().unwrap_or_else(|e| e.into_inner()).len(), + ) .finish() } } @@ -526,6 +539,20 @@ mod tests { assert!(buffer.capacity() >= 2048); } + #[tokio::test] + async fn test_acquire_buffer_after_shutdown_is_unpooled() { + let pool = BytesPool::new_tiered(); + pool.small_pool.semaphore.close(); + + let buffer = pool.acquire_buffer(2048).await; + + assert!(buffer.tier.is_none()); + assert!(buffer._permit.is_none()); + assert!(buffer.capacity() >= pool.small_pool.buffer_size); + drop(buffer); + assert_eq!(pool.available_buffers(), 0); + } + #[tokio::test] async fn test_tier_selection() { let pool = BytesPool::new_tiered(); diff --git a/crates/signer/src/lib.rs b/crates/signer/src/lib.rs index 582294d07..d802b0f32 100644 --- a/crates/signer/src/lib.rs +++ b/crates/signer/src/lib.rs @@ -20,8 +20,12 @@ pub mod request_signature_v4; pub mod utils; pub use request_signature_streaming::streaming_sign_v4; +pub use request_signature_streaming::try_streaming_sign_v4; +pub use request_signature_v2::SignV2Error; pub use request_signature_v2::pre_sign_v2; pub use request_signature_v2::sign_v2; +pub use request_signature_v2::try_pre_sign_v2; +pub use request_signature_v2::try_sign_v2; pub use request_signature_v4::SignV4Error; pub use request_signature_v4::pre_sign_v4; pub use request_signature_v4::sign_v4; diff --git a/crates/signer/src/request_signature_streaming.rs b/crates/signer/src/request_signature_streaming.rs index 85469f592..a3db101ca 100644 --- a/crates/signer/src/request_signature_streaming.rs +++ b/crates/signer/src/request_signature_streaming.rs @@ -14,8 +14,9 @@ use http::{HeaderMap, HeaderValue, request}; use time::{OffsetDateTime, macros::format_description}; +use tracing::warn; -use super::request_signature_v4::{SERVICE_TYPE_S3, get_scope, get_signature, get_signing_key}; +use super::request_signature_v4::{SERVICE_TYPE_S3, SignV4Error, get_scope, get_signature, get_signing_key}; use rustfs_utils::hash::EMPTY_STRING_SHA256_HASH; use s3s::Body; @@ -38,67 +39,194 @@ const _TRAILER_SIGNATURE: &str = "x-amz-trailer-signature"; // m // }); +#[derive(Debug)] +struct StreamingSignFailure { + request: request::Request, + error: SignV4Error, +} + +type StreamingSignOutcome = std::result::Result, Box>; + +fn streaming_fail(request: request::Request, error: SignV4Error) -> StreamingSignOutcome { + Err(Box::new(StreamingSignFailure { request, error })) +} + #[allow(dead_code)] -fn build_chunk_string_to_sign(t: OffsetDateTime, region: &str, previous_sig: &str, chunk_check_sum: &str) -> String { +fn try_build_chunk_string_to_sign( + t: OffsetDateTime, + region: &str, + previous_sig: &str, + chunk_check_sum: &str, +) -> Result { let mut string_to_sign_parts = >::new(); string_to_sign_parts.push(STREAMING_PAYLOAD_HDR.to_string()); let format = format_description!("[year][month][day]T[hour][minute][second]Z"); - string_to_sign_parts.push(t.format(&format).unwrap()); + string_to_sign_parts.push( + t.format(&format) + .map_err(|err| SignV4Error::TimeFormat { reason: err.to_string() })?, + ); string_to_sign_parts.push(get_scope(region, t, SERVICE_TYPE_S3)); string_to_sign_parts.push(previous_sig.to_string()); string_to_sign_parts.push(EMPTY_STRING_SHA256_HASH.to_string()); string_to_sign_parts.push(chunk_check_sum.to_string()); - string_to_sign_parts.join("\n") + Ok(string_to_sign_parts.join("\n")) } -fn _build_chunk_signature( +fn _try_build_chunk_signature( chunk_check_sum: &str, req_time: OffsetDateTime, region: &str, previous_signature: &str, secret_access_key: &str, -) -> String { - let chunk_string_to_sign = build_chunk_string_to_sign(req_time, region, previous_signature, chunk_check_sum); +) -> Result { + let chunk_string_to_sign = try_build_chunk_string_to_sign(req_time, region, previous_signature, chunk_check_sum)?; let signing_key = get_signing_key(secret_access_key, region, req_time, SERVICE_TYPE_S3); - get_signature(signing_key, &chunk_string_to_sign) + Ok(get_signature(signing_key, &chunk_string_to_sign)) } #[allow(clippy::too_many_arguments)] -pub fn streaming_sign_v4( +fn streaming_sign_v4_inner( mut req: request::Request, _access_key_id: &str, _secret_access_key: &str, session_token: &str, _region: &str, data_len: i64, - req_time: OffsetDateTime, /*, sh256: md5simd::Hasher*/ + req_time: OffsetDateTime, trailer: HeaderMap, -) -> request::Request { +) -> StreamingSignOutcome { let headers = req.headers_mut(); if trailer.is_empty() { - headers.append("X-Amz-Content-Sha256", HeaderValue::from_str(STREAMING_SIGN_ALGORITHM).expect("err")); + let value = match HeaderValue::from_str(STREAMING_SIGN_ALGORITHM) { + Ok(v) => v, + Err(err) => { + return streaming_fail( + req, + SignV4Error::HeaderValueParse { + name: "X-Amz-Content-Sha256".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.append("X-Amz-Content-Sha256", value); } else { - headers.append( - "X-Amz-Content-Sha256", - HeaderValue::from_str(STREAMING_SIGN_TRAILER_ALGORITHM).expect("err"), - ); + let trailer_algo = match HeaderValue::from_str(STREAMING_SIGN_TRAILER_ALGORITHM) { + Ok(v) => v, + Err(err) => { + return streaming_fail( + req, + SignV4Error::HeaderValueParse { + name: "X-Amz-Content-Sha256".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.append("X-Amz-Content-Sha256", trailer_algo); for (k, _) in &trailer { - headers.append("X-Amz-Trailer", k.as_str().to_lowercase().parse().unwrap()); + let parsed = match k.as_str().to_lowercase().parse::() { + Ok(v) => v, + Err(err) => { + return streaming_fail( + req, + SignV4Error::HeaderValueParse { + name: "X-Amz-Trailer".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.append("X-Amz-Trailer", parsed); } - let chunked_value = HeaderValue::from_str(&["aws-chunked"].join(",")).expect("err"); - headers.insert(http::header::TRANSFER_ENCODING, chunked_value); + headers.insert(http::header::TRANSFER_ENCODING, HeaderValue::from_static("aws-chunked")); } if !session_token.is_empty() { - headers.insert("X-Amz-Security-Token", HeaderValue::from_str(session_token).expect("err")); + let token_value = match HeaderValue::from_str(session_token) { + Ok(v) => v, + Err(err) => { + return streaming_fail( + req, + SignV4Error::HeaderValueParse { + name: "X-Amz-Security-Token".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.insert("X-Amz-Security-Token", token_value); } let format = format_description!("[year]-[month]-[day]T[hour]:[minute]:[second].[subsecond]Z"); - headers.insert("X-Amz-Date", HeaderValue::from_str(&req_time.format(&format).unwrap()).expect("err")); + let date_str = match req_time.format(&format) { + Ok(v) => v, + Err(err) => return streaming_fail(req, SignV4Error::TimeFormat { reason: err.to_string() }), + }; + let date_value = match HeaderValue::from_str(&date_str) { + Ok(v) => v, + Err(err) => { + return streaming_fail( + req, + SignV4Error::HeaderValueParse { + name: "X-Amz-Date".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.insert("X-Amz-Date", date_value); //req.content_length = 100; - headers.insert("x-amz-decoded-content-length", format!("{data_len:010}").parse().unwrap()); + let decoded_len = match format!("{data_len:010}").parse::() { + Ok(v) => v, + Err(err) => { + return streaming_fail( + req, + SignV4Error::HeaderValueParse { + name: "x-amz-decoded-content-length".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.insert("x-amz-decoded-content-length", decoded_len); - req + Ok(req) +} + +#[allow(clippy::too_many_arguments)] +pub fn try_streaming_sign_v4( + req: request::Request, + access_key_id: &str, + secret_access_key: &str, + session_token: &str, + region: &str, + data_len: i64, + req_time: OffsetDateTime, + trailer: HeaderMap, +) -> Result, SignV4Error> { + streaming_sign_v4_inner(req, access_key_id, secret_access_key, session_token, region, data_len, req_time, trailer) + .map_err(|f| f.error) +} + +#[allow(clippy::too_many_arguments)] +pub fn streaming_sign_v4( + req: request::Request, + access_key_id: &str, + secret_access_key: &str, + session_token: &str, + region: &str, + data_len: i64, + req_time: OffsetDateTime, + trailer: HeaderMap, +) -> request::Request { + match streaming_sign_v4_inner(req, access_key_id, secret_access_key, session_token, region, data_len, req_time, trailer) { + Ok(request) => request, + Err(failure) => { + warn!(error = %failure.error, "failed to sign streaming v4 request"); + failure.request + } + } } diff --git a/crates/signer/src/request_signature_v2.rs b/crates/signer/src/request_signature_v2.rs index 35d7bdd9b..558dfe45b 100644 --- a/crates/signer/src/request_signature_v2.rs +++ b/crates/signer/src/request_signature_v2.rs @@ -18,46 +18,105 @@ use hyper::Uri; use std::collections::HashMap; use std::fmt::Write; use time::{OffsetDateTime, format_description}; +use tracing::warn; -use super::utils::get_host_addr; +use http::HeaderValue; + +use super::utils::{HostAddrError, try_get_host_addr}; use rustfs_utils::crypto::{hex, hmac_sha1}; use s3s::Body; const _SIGN_V4_ALGORITHM: &str = "AWS4-HMAC-SHA256"; const SIGN_V2_ALGORITHM: &str = "AWS"; +#[derive(Debug, thiserror::Error)] +pub enum SignV2Error { + #[error("invalid UTF-8 header value for `{name}`")] + InvalidHeaderValue { name: String }, + #[error("failed to format signing timestamp: {reason}")] + TimeFormat { reason: String }, + #[error("failed to build signing timestamp: {reason}")] + TimeComponent { reason: String }, + #[error("failed to encode query parameters: {reason}")] + QueryEncode { reason: String }, + #[error("failed to parse uri: {reason}")] + InvalidUri { reason: String }, + #[error("failed to build uri from parts: {reason}")] + InvalidUriParts { reason: String }, + #[error("failed to convert canonical headers to UTF-8: {reason}")] + CanonicalUtf8 { reason: String }, + #[error("failed to parse header value for `{name}`: {reason}")] + HeaderValueParse { name: String, reason: String }, + #[error("failed to resolve host address: {0}")] + HostAddr(#[from] HostAddrError), +} + +#[derive(Debug)] +struct SignV2Failure { + request: request::Request, + error: SignV2Error, +} + +type SignV2Outcome = std::result::Result, Box>; + +fn sign_v2_fail(request: request::Request, error: SignV2Error) -> SignV2Outcome { + Err(Box::new(SignV2Failure { request, error })) +} + fn encode_url2path(req: &request::Request, _virtual_host: bool) -> String { req.uri().path().to_string() } -pub fn pre_sign_v2( +fn pre_sign_v2_inner( mut req: request::Request, access_key_id: &str, secret_access_key: &str, expires: i64, virtual_host: bool, -) -> request::Request { +) -> SignV2Outcome { if access_key_id.is_empty() || secret_access_key.is_empty() { - return req; + return Ok(req); } let d = OffsetDateTime::now_utc(); - let d = d.replace_time(time::Time::from_hms(0, 0, 0).unwrap()); + let d = match time::Time::from_hms(0, 0, 0) { + Ok(midnight) => d.replace_time(midnight), + Err(err) => return sign_v2_fail(req, SignV2Error::TimeComponent { reason: err.to_string() }), + }; let epoch_expires = d.unix_timestamp() + expires; let headers = req.headers_mut(); let expires_str = headers.get("Expires"); if expires_str.is_none() { - headers.insert("Expires", format!("{epoch_expires:010}").parse().unwrap()); + let expires_value = match format!("{epoch_expires:010}").parse::() { + Ok(v) => v, + Err(err) => { + return sign_v2_fail( + req, + SignV2Error::HeaderValueParse { + name: "Expires".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.insert("Expires", expires_value); } - let string_to_sign = pre_string_to_sign_v2(&req, virtual_host); + let string_to_sign = match try_pre_string_to_sign_v2(&req, virtual_host) { + Ok(v) => v, + Err(err) => return sign_v2_fail(req, err), + }; let signature = hex(hmac_sha1(secret_access_key, string_to_sign)); let query_source = req.uri().query().unwrap_or(""); let result = serde_urlencoded::from_str::>(query_source); let mut query = result.unwrap_or_default(); - if get_host_addr(&req).contains(".storage.googleapis.com") { + let host_addr = match try_get_host_addr(&req) { + Ok(v) => v, + Err(err) => return sign_v2_fail(req, SignV2Error::HostAddr(err)), + }; + if host_addr.contains(".storage.googleapis.com") { query.insert("GoogleAccessId".to_string(), access_key_id.to_string()); } else { query.insert("AWSAccessKeyId".to_string(), access_key_id.to_string()); @@ -67,43 +126,97 @@ pub fn pre_sign_v2( let uri = req.uri().clone(); let mut parts = req.uri().clone().into_parts(); - parts.path_and_query = Some( - format!("{}?{}&Signature={}", uri.path(), serde_urlencoded::to_string(&query).unwrap(), signature) - .parse() - .unwrap(), - ); + let query_str = match serde_urlencoded::to_string(&query) { + Ok(v) => v, + Err(err) => return sign_v2_fail(req, SignV2Error::QueryEncode { reason: err.to_string() }), + }; + parts.path_and_query = Some(match format!("{}?{}&Signature={}", uri.path(), query_str, signature).parse() { + Ok(v) => v, + Err(err) => return sign_v2_fail(req, SignV2Error::InvalidUri { reason: err.to_string() }), + }); - *req.uri_mut() = Uri::from_parts(parts).unwrap(); + *req.uri_mut() = match Uri::from_parts(parts) { + Ok(v) => v, + Err(err) => return sign_v2_fail(req, SignV2Error::InvalidUriParts { reason: err.to_string() }), + }; - req + Ok(req) +} + +pub fn try_pre_sign_v2( + req: request::Request, + access_key_id: &str, + secret_access_key: &str, + expires: i64, + virtual_host: bool, +) -> Result, SignV2Error> { + pre_sign_v2_inner(req, access_key_id, secret_access_key, expires, virtual_host).map_err(|f| f.error) +} + +pub fn pre_sign_v2( + req: request::Request, + access_key_id: &str, + secret_access_key: &str, + expires: i64, + virtual_host: bool, +) -> request::Request { + match pre_sign_v2_inner(req, access_key_id, secret_access_key, expires, virtual_host) { + Ok(request) => request, + Err(failure) => { + warn!(error = %failure.error, "failed to presign v2 request"); + failure.request + } + } } fn _post_pre_sign_signature_v2(policy_base64: &str, secret_access_key: &str) -> String { hex(hmac_sha1(secret_access_key, policy_base64)) } -pub fn sign_v2( +fn sign_v2_inner( mut req: request::Request, _content_len: i64, access_key_id: &str, secret_access_key: &str, virtual_host: bool, -) -> request::Request { +) -> SignV2Outcome { if access_key_id.is_empty() || secret_access_key.is_empty() { - return req; + return Ok(req); } let d = OffsetDateTime::now_utc(); - let d2 = d.replace_time(time::Time::from_hms(0, 0, 0).unwrap()); + let d2 = match time::Time::from_hms(0, 0, 0) { + Ok(midnight) => d.replace_time(midnight), + Err(err) => return sign_v2_fail(req, SignV2Error::TimeComponent { reason: err.to_string() }), + }; { let headers = req.headers_mut(); let need_default_date = headers.get("Date").and_then(|v| v.to_str().ok()).is_none_or(|v| v.is_empty()); if need_default_date { - headers.insert("Date", d2.format(&format_description::well_known::Rfc2822).unwrap().parse().unwrap()); + let date_str = match d2.format(&format_description::well_known::Rfc2822) { + Ok(v) => v, + Err(err) => return sign_v2_fail(req, SignV2Error::TimeFormat { reason: err.to_string() }), + }; + let date_value = match date_str.parse::() { + Ok(v) => v, + Err(err) => { + return sign_v2_fail( + req, + SignV2Error::HeaderValueParse { + name: "Date".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.insert("Date", date_value); } } - let string_to_sign = string_to_sign_v2(&req, virtual_host); + let string_to_sign = match try_string_to_sign_v2(&req, virtual_host) { + Ok(v) => v, + Err(err) => return sign_v2_fail(req, err), + }; let headers = req.headers_mut(); let auth_header = format!("{SIGN_V2_ALGORITHM} {access_key_id}:"); @@ -113,17 +226,55 @@ pub fn sign_v2( base64_simd::URL_SAFE_NO_PAD.encode_to_string(hmac_sha1(secret_access_key, string_to_sign)) ); - headers.insert("Authorization", auth_header.parse().unwrap()); + let auth_value = match auth_header.parse::() { + Ok(v) => v, + Err(err) => { + return sign_v2_fail( + req, + SignV2Error::HeaderValueParse { + name: "Authorization".to_string(), + reason: err.to_string(), + }, + ); + } + }; + headers.insert("Authorization", auth_value); - req + Ok(req) } -fn pre_string_to_sign_v2(req: &request::Request, virtual_host: bool) -> String { +pub fn try_sign_v2( + req: request::Request, + content_len: i64, + access_key_id: &str, + secret_access_key: &str, + virtual_host: bool, +) -> Result, SignV2Error> { + sign_v2_inner(req, content_len, access_key_id, secret_access_key, virtual_host).map_err(|f| f.error) +} + +pub fn sign_v2( + req: request::Request, + content_len: i64, + access_key_id: &str, + secret_access_key: &str, + virtual_host: bool, +) -> request::Request { + match sign_v2_inner(req, content_len, access_key_id, secret_access_key, virtual_host) { + Ok(request) => request, + Err(failure) => { + warn!(error = %failure.error, "failed to sign v2 request"); + failure.request + } + } +} + +fn try_pre_string_to_sign_v2(req: &request::Request, virtual_host: bool) -> Result { let mut buf = BytesMut::new(); write_pre_sign_v2_headers(&mut buf, req); write_canonicalized_headers(&mut buf, req); write_canonicalized_resource(&mut buf, req, virtual_host); - String::from_utf8(buf.to_vec()).unwrap() + String::from_utf8(buf.to_vec()).map_err(|err| SignV2Error::CanonicalUtf8 { reason: err.to_string() }) } fn write_pre_sign_v2_headers(buf: &mut BytesMut, req: &request::Request) { @@ -137,12 +288,12 @@ fn write_pre_sign_v2_headers(buf: &mut BytesMut, req: &request::Request) { let _ = buf.write_char('\n'); } -fn string_to_sign_v2(req: &request::Request, virtual_host: bool) -> String { +fn try_string_to_sign_v2(req: &request::Request, virtual_host: bool) -> Result { let mut buf = BytesMut::new(); write_sign_v2_headers(&mut buf, req); write_canonicalized_headers(&mut buf, req); write_canonicalized_resource(&mut buf, req, virtual_host); - String::from_utf8(buf.to_vec()).unwrap() + String::from_utf8(buf.to_vec()).map_err(|err| SignV2Error::CanonicalUtf8 { reason: err.to_string() }) } fn write_sign_v2_headers(buf: &mut BytesMut, req: &request::Request) { @@ -309,7 +460,7 @@ mod tests { .insert("host", "examplebucket.s3.amazonaws.com".parse().unwrap()); let req = sign_v2(req, 0, "AKIAEXAMPLE", "SECRET", false); - let expected_string_to_sign = string_to_sign_v2(&req, false); + let expected_string_to_sign = try_string_to_sign_v2(&req, false).expect("string to sign should build"); let expected_signature = base64_simd::URL_SAFE_NO_PAD.encode_to_string(hmac_sha1("SECRET", expected_string_to_sign)); assert_eq!( diff --git a/crates/signer/src/utils.rs b/crates/signer/src/utils.rs index 7a8710d83..0141403e1 100644 --- a/crates/signer/src/utils.rs +++ b/crates/signer/src/utils.rs @@ -46,7 +46,15 @@ pub fn try_get_host_addr(req: &request::Request) -> Result) -> String { - try_get_host_addr(req).unwrap() + match try_get_host_addr(req) { + Ok(host) => host, + Err(HostAddrError::MissingUriHost) => match req.headers().get("host").map(|host| host.to_str()) { + Some(Ok(host)) => host.to_string(), + Some(Err(_)) => panic!("failed to resolve request host: invalid UTF-8 header value for `host`"), + None => panic!("failed to resolve request host: request uri has no host"), + }, + Err(err) => panic!("failed to resolve request host: {err}"), + } } pub fn sign_v4_trim_all(input: &str) -> String { @@ -63,7 +71,7 @@ where #[cfg(test)] mod tests { - use super::{HostAddrError, try_get_host_addr}; + use super::{HostAddrError, get_host_addr, try_get_host_addr}; use http::HeaderValue; use http::request; use s3s::Body; @@ -83,6 +91,30 @@ mod tests { assert_eq!(host, "proxy.internal:9443"); } + #[test] + fn get_host_addr_preserves_legacy_string_api() { + let req = request::Request::builder() + .method(http::Method::GET) + .uri("https://bucket.example.com:9443/object") + .body(Body::empty()) + .expect("request should build"); + + assert_eq!(get_host_addr(&req), "bucket.example.com:9443"); + } + + #[test] + fn get_host_addr_uses_host_header_for_relative_uri() { + let mut req = request::Request::builder() + .method(http::Method::GET) + .uri("/object") + .body(Body::empty()) + .expect("request should build"); + req.headers_mut() + .insert("host", HeaderValue::from_static("bucket.example.com")); + + assert_eq!(get_host_addr(&req), "bucket.example.com"); + } + #[test] fn try_get_host_addr_rejects_non_utf8_host_header_value() { let mut req = request::Request::builder()