Merge branch 'main' into feat/kms-vault-transit2

This commit is contained in:
安正超
2026-04-07 08:36:09 +08:00
committed by GitHub
162 changed files with 23812 additions and 5695 deletions
+2
View File
@@ -47,6 +47,7 @@ webdav = ["rustfs-protocols/webdav"]
license = []
direct-io = [] # Aligned direct I/O reader support (uses aligned pread, does not set O_DIRECT)
io-scheduler-debug = [] # Enable debug information in I/O scheduler
tracing-chunk-debug = [] # Enable per-chunk tracing in data plane (high noise, for debugging only)
full = ["metrics-gpu", "ftps", "swift", "webdav", "direct-io"]
manual-test-runners = []
@@ -85,6 +86,7 @@ rustfs-utils = { workspace = true, features = ["full"] }
rustfs-zip = { workspace = true }
rustfs-io-core = { workspace = true }
rustfs-io-metrics = { workspace = true }
rustfs-object-io = { workspace = true }
rustfs-concurrency = { workspace = true }
rustfs-scanner = { workspace = true }
+895
View File
@@ -0,0 +1,895 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use crate::admin::router::{AdminOperation, Operation, S3Router};
use crate::auth::{check_key_valid, get_session_token};
use crate::server::ADMIN_PREFIX;
use futures::stream::{FuturesUnordered, StreamExt};
use hashbrown::HashSet as HbHashSet;
use http::{HeaderMap, StatusCode};
use hyper::Method;
use matchit::Params;
use rustfs_audit::{audit_system, start_audit_system as start_global_audit_system, system::AuditSystemState};
use rustfs_config::audit::{AUDIT_MQTT_KEYS, AUDIT_MQTT_SUB_SYS, AUDIT_ROUTE_PREFIX, AUDIT_WEBHOOK_KEYS, AUDIT_WEBHOOK_SUB_SYS};
use rustfs_config::{DEFAULT_DELIMITER, ENABLE_KEY, ENV_PREFIX, EnableState, MAX_ADMIN_REQUEST_BODY_SIZE};
use rustfs_ecstore::config::Config;
use rustfs_targets::check_mqtt_broker_available;
use s3s::{Body, S3Request, S3Response, S3Result, header::CONTENT_TYPE, s3_error};
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
use std::future::Future;
use std::io::{Error, ErrorKind};
use std::path::Path;
use std::sync::Arc;
use tokio::sync::Semaphore;
use tokio::time::{Duration, sleep, timeout};
use tracing::{Span, warn};
use url::Url;
pub fn register_audit_target_route(r: &mut S3Router<AdminOperation>) -> std::io::Result<()> {
r.insert(
Method::GET,
format!("{}{}", ADMIN_PREFIX, "/v3/audit/target/list").as_str(),
AdminOperation(&ListAuditTargets {}),
)?;
r.insert(
Method::PUT,
format!("{}{}", ADMIN_PREFIX, "/v3/audit/target/{target_type}/{target_name}").as_str(),
AdminOperation(&AuditTargetConfig {}),
)?;
r.insert(
Method::DELETE,
format!("{}{}", ADMIN_PREFIX, "/v3/audit/target/{target_type}/{target_name}/reset").as_str(),
AdminOperation(&RemoveAuditTarget {}),
)?;
Ok(())
}
#[derive(Debug, Deserialize)]
pub struct KeyValue {
pub key: String,
pub value: String,
}
#[derive(Debug, Deserialize)]
pub struct AuditTargetBody {
pub key_values: Vec<KeyValue>,
}
#[derive(Serialize, Debug)]
struct AuditEndpoint {
account_id: String,
service: String,
status: String,
source: AuditEndpointSource,
}
#[derive(Serialize, Debug)]
struct AuditEndpointsResponse {
audit_endpoints: Vec<AuditEndpoint>,
}
type EndpointKey = (String, String);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
enum AuditEndpointSource {
Config,
Env,
Mixed,
Runtime,
}
fn normalized_endpoint_key(account_id: &str, service: &str) -> EndpointKey {
(account_id.to_lowercase(), service.to_string())
}
async fn check_permissions(req: &S3Request<Body>) -> S3Result<()> {
let Some(input_cred) = &req.credentials else {
return Err(s3_error!(InvalidRequest, "credentials not found"));
};
check_key_valid(get_session_token(&req.uri, &req.headers).unwrap_or_default(), &input_cred.access_key).await?;
Ok(())
}
fn build_response(status: StatusCode, body: Body, request_id: Option<&http::HeaderValue>) -> S3Response<(StatusCode, Body)> {
let mut header = HeaderMap::new();
header.insert(CONTENT_TYPE, "application/json".parse().unwrap());
if let Some(v) = request_id {
header.insert("x-request-id", v.clone());
}
S3Response::with_headers((status, body), header)
}
async fn retry_with_backoff<F, Fut, T>(mut operation: F, max_attempts: usize, base_delay: Duration) -> Result<T, Error>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<T, Error>>,
{
let mut attempts = 0;
let mut delay = base_delay;
let mut last_err = None;
while attempts < max_attempts {
match operation().await {
Ok(result) => return Ok(result),
Err(e) => {
last_err = Some(e);
attempts += 1;
if attempts < max_attempts {
sleep(delay).await;
delay = delay.saturating_mul(2);
}
}
}
}
Err(last_err.unwrap_or_else(|| Error::other("retry_with_backoff: unknown error")))
}
async fn validate_queue_dir(queue_dir: &str) -> S3Result<()> {
if !queue_dir.is_empty() {
if !Path::new(queue_dir).is_absolute() {
return Err(s3_error!(InvalidArgument, "queue_dir must be absolute path"));
}
retry_with_backoff(
|| async { tokio::fs::metadata(queue_dir).await.map(|_| ()) },
3,
Duration::from_millis(100),
)
.await
.map_err(|e| match e.kind() {
ErrorKind::NotFound => s3_error!(InvalidArgument, "queue_dir does not exist"),
ErrorKind::PermissionDenied => s3_error!(InvalidArgument, "queue_dir exists but permission denied"),
_ => s3_error!(InvalidArgument, "failed to access queue_dir: {}", e),
})?;
}
Ok(())
}
fn config_enable_is_on(value: &str) -> bool {
matches!(value.trim().to_ascii_lowercase().as_str(), "on" | "true" | "yes" | "1")
}
fn has_any_audit_targets(config: &Config) -> bool {
for subsystem in [AUDIT_WEBHOOK_SUB_SYS, AUDIT_MQTT_SUB_SYS] {
let Some(targets) = config.0.get(subsystem) else {
continue;
};
if targets.keys().any(|key| key != DEFAULT_DELIMITER) {
return true;
}
}
false
}
fn collect_configured_audit_endpoint_keys(config: &Config) -> Vec<EndpointKey> {
let mut endpoints = Vec::new();
for (subsystem, service) in [(AUDIT_WEBHOOK_SUB_SYS, "webhook"), (AUDIT_MQTT_SUB_SYS, "mqtt")] {
let Some(targets) = config.0.get(subsystem) else {
continue;
};
for (target_name, kvs) in targets {
if target_name == DEFAULT_DELIMITER {
continue;
}
let enabled = kvs.lookup(ENABLE_KEY).as_deref().map(config_enable_is_on).unwrap_or(false);
if enabled {
endpoints.push((target_name.clone(), service.to_string()));
}
}
}
endpoints
}
fn collect_config_entry_keys(config: &Config) -> HbHashSet<EndpointKey> {
let mut endpoints = HbHashSet::new();
for (subsystem, service) in [(AUDIT_WEBHOOK_SUB_SYS, "webhook"), (AUDIT_MQTT_SUB_SYS, "mqtt")] {
let Some(targets) = config.0.get(subsystem) else {
continue;
};
for target_name in targets.keys() {
if target_name == DEFAULT_DELIMITER {
continue;
}
endpoints.insert(normalized_endpoint_key(target_name, service));
}
}
endpoints
}
fn collect_env_endpoint_keys() -> HbHashSet<EndpointKey> {
let mut endpoints = HbHashSet::new();
for (service, valid_keys) in [("webhook", AUDIT_WEBHOOK_KEYS), ("mqtt", AUDIT_MQTT_KEYS)] {
let env_prefix = format!("{ENV_PREFIX}{AUDIT_ROUTE_PREFIX}{service}{DEFAULT_DELIMITER}").to_uppercase();
for (key, _value) in std::env::vars() {
let Some(rest) = key.strip_prefix(&env_prefix) else {
continue;
};
let mut parts = rest.rsplitn(2, DEFAULT_DELIMITER);
let instance_id_part = parts.next().unwrap_or(DEFAULT_DELIMITER);
let field_name_part = parts.next();
let (field_name, instance_id) = match field_name_part {
Some(field) => (field.to_lowercase(), instance_id_part.to_lowercase()),
None => (instance_id_part.to_lowercase(), DEFAULT_DELIMITER.to_string()),
};
if instance_id == DEFAULT_DELIMITER || instance_id.is_empty() {
continue;
}
if valid_keys.contains(&field_name.as_str()) {
endpoints.insert(normalized_endpoint_key(&instance_id, service));
}
}
}
endpoints
}
fn classify_audit_endpoint_source(
config_targets: &HbHashSet<EndpointKey>,
env_targets: &HbHashSet<EndpointKey>,
key: &EndpointKey,
) -> AuditEndpointSource {
match (config_targets.contains(key), env_targets.contains(key)) {
(true, true) => AuditEndpointSource::Mixed,
(true, false) => AuditEndpointSource::Config,
(false, true) => AuditEndpointSource::Env,
(false, false) => AuditEndpointSource::Runtime,
}
}
fn audit_endpoint_source(config: &Config, target_type: &str, target_name: &str) -> AuditEndpointSource {
let config_targets = collect_config_entry_keys(config);
let env_targets = collect_env_endpoint_keys();
let service = match target_type {
AUDIT_WEBHOOK_SUB_SYS => "webhook",
AUDIT_MQTT_SUB_SYS => "mqtt",
_ => "",
};
let key = normalized_endpoint_key(target_name, service);
classify_audit_endpoint_source(&config_targets, &env_targets, &key)
}
fn audit_target_mutation_block_reason(config: &Config, target_type: &str, target_name: &str) -> Option<String> {
match audit_endpoint_source(config, target_type, target_name) {
AuditEndpointSource::Env => Some(format!(
"audit target '{}' is managed by environment variables and cannot be modified from the console",
target_name
)),
AuditEndpointSource::Mixed => Some(format!(
"audit target '{}' is configured by both persisted config and environment variables; remove the environment variables first",
target_name
)),
AuditEndpointSource::Config | AuditEndpointSource::Runtime => None,
}
}
fn merge_audit_endpoints(config: &Config, runtime_statuses: HashMap<EndpointKey, String>) -> Vec<AuditEndpoint> {
let mut audit_endpoints = Vec::new();
let mut seen = HashSet::new();
let configured_keys = collect_configured_audit_endpoint_keys(config);
let config_targets = collect_config_entry_keys(config);
let env_targets = collect_env_endpoint_keys();
let mut normalized_runtime_statuses: HashMap<EndpointKey, (String, String, String)> = HashMap::new();
for ((account_id, service), status) in runtime_statuses {
let normalized = normalized_endpoint_key(&account_id, &service);
normalized_runtime_statuses
.entry(normalized)
.or_insert((account_id, service, status));
}
for key in configured_keys {
let normalized = normalized_endpoint_key(&key.0, &key.1);
if !seen.insert(normalized.clone()) {
continue;
}
let status = normalized_runtime_statuses
.remove(&normalized)
.map(|(_, _, status)| status)
.unwrap_or_else(|| "offline".to_string());
let source = classify_audit_endpoint_source(&config_targets, &env_targets, &normalized);
audit_endpoints.push(AuditEndpoint {
account_id: key.0,
service: key.1,
status,
source,
});
}
for (normalized, (account_id, service, status)) in normalized_runtime_statuses {
if seen.insert(normalized.clone()) {
audit_endpoints.push(AuditEndpoint {
account_id,
service,
status,
source: classify_audit_endpoint_source(&config_targets, &env_targets, &normalized),
});
}
}
for key in &env_targets {
if !seen.insert(key.clone()) {
continue;
}
audit_endpoints.push(AuditEndpoint {
account_id: key.0.clone(),
service: key.1.clone(),
status: "offline".to_string(),
source: classify_audit_endpoint_source(&config_targets, &env_targets, key),
});
}
audit_endpoints.sort_by(|a, b| a.service.cmp(&b.service).then_with(|| a.account_id.cmp(&b.account_id)));
audit_endpoints
}
fn collect_validated_key_values(
key_values: &[KeyValue],
allowed_keys: &HashSet<&str>,
target_type: &str,
) -> S3Result<HashMap<String, String>> {
let mut kv_map = HashMap::new();
let mut seen = HashSet::new();
for kv in key_values {
if !allowed_keys.contains(kv.key.as_str()) {
return Err(s3_error!(
InvalidArgument,
"key '{}' not allowed for audit target type '{}'",
kv.key,
target_type
));
}
if !seen.insert(kv.key.as_str()) {
return Err(s3_error!(InvalidArgument, "duplicate key '{}' in request body", kv.key));
}
kv_map.insert(kv.key.clone(), kv.value.clone());
}
Ok(kv_map)
}
fn extract_target_params<'a>(params: &'a Params<'_, '_>) -> S3Result<(&'a str, &'a str)> {
let target_type = params
.get("target_type")
.ok_or_else(|| s3_error!(InvalidArgument, "missing required parameter: 'target_type'"))?;
if target_type != AUDIT_WEBHOOK_SUB_SYS && target_type != AUDIT_MQTT_SUB_SYS {
return Err(s3_error!(InvalidArgument, "unsupported audit target type: '{}'", target_type));
}
let target_name = params
.get("target_name")
.ok_or_else(|| s3_error!(InvalidArgument, "missing required parameter: 'target_name'"))?;
Ok((target_type, target_name))
}
async fn load_server_config_from_store() -> S3Result<Config> {
let Some(store) = rustfs_ecstore::global::new_object_layer_fn() else {
return Ok(Config::new());
};
rustfs_ecstore::config::com::read_config_without_migrate(store)
.await
.map_err(|e| s3_error!(InternalError, "failed to read server config: {}", e))
}
async fn apply_audit_runtime_config(config: Config) -> S3Result<()> {
let has_targets = has_any_audit_targets(&config);
if let Some(system) = audit_system() {
match system.get_state().await {
AuditSystemState::Running | AuditSystemState::Paused | AuditSystemState::Starting => {
if has_targets {
system
.reload_config(config)
.await
.map_err(|e| s3_error!(InternalError, "failed to reload audit config: {}", e))?;
} else {
system
.close()
.await
.map_err(|e| s3_error!(InternalError, "failed to stop audit system: {}", e))?;
}
}
AuditSystemState::Stopped | AuditSystemState::Stopping => {
if has_targets {
system
.start(config)
.await
.map_err(|e| s3_error!(InternalError, "failed to start audit system: {}", e))?;
}
}
}
} else if has_targets {
start_global_audit_system(config)
.await
.map_err(|e| s3_error!(InternalError, "failed to start audit system: {}", e))?;
}
Ok(())
}
async fn update_audit_config_and_reload<F>(mut modifier: F) -> S3Result<()>
where
F: FnMut(&mut Config) -> bool,
{
let Some(store) = rustfs_ecstore::global::new_object_layer_fn() else {
return Err(s3_error!(InternalError, "server storage not initialized"));
};
let mut config = rustfs_ecstore::config::com::read_config_without_migrate(store.clone())
.await
.map_err(|e| s3_error!(InternalError, "failed to read server config: {}", e))?;
if !modifier(&mut config) {
return Ok(());
}
rustfs_ecstore::config::com::save_server_config(store, &config)
.await
.map_err(|e| s3_error!(InternalError, "failed to save audit config: {}", e))?;
apply_audit_runtime_config(config).await
}
pub struct AuditTargetConfig {}
#[async_trait::async_trait]
impl Operation for AuditTargetConfig {
async fn call(&self, req: S3Request<Body>, params: Params<'_, '_>) -> S3Result<S3Response<(StatusCode, Body)>> {
let span = Span::current();
let _enter = span.enter();
let (target_type, target_name) = extract_target_params(&params)?;
check_permissions(&req).await?;
let config_snapshot = load_server_config_from_store().await?;
if let Some(reason) = audit_target_mutation_block_reason(&config_snapshot, target_type, target_name) {
return Err(s3_error!(InvalidRequest, "{reason}"));
}
let mut input = req.input;
let body_bytes = input.store_all_limited(MAX_ADMIN_REQUEST_BODY_SIZE).await.map_err(|e| {
warn!("failed to read request body: {:?}", e);
s3_error!(InvalidRequest, "failed to read request body")
})?;
let audit_body: AuditTargetBody = serde_json::from_slice(&body_bytes)
.map_err(|e| s3_error!(InvalidArgument, "invalid json body for audit target config: {}", e))?;
let allowed_keys: HashSet<&str> = match target_type {
AUDIT_WEBHOOK_SUB_SYS => AUDIT_WEBHOOK_KEYS.iter().cloned().collect(),
AUDIT_MQTT_SUB_SYS => AUDIT_MQTT_KEYS.iter().cloned().collect(),
_ => unreachable!(),
};
let kv_map = collect_validated_key_values(&audit_body.key_values, &allowed_keys, target_type)?;
if target_type == AUDIT_WEBHOOK_SUB_SYS {
let endpoint = kv_map
.get("endpoint")
.map(String::as_str)
.ok_or_else(|| s3_error!(InvalidArgument, "endpoint is required"))?;
let parsed_endpoint = Url::parse(endpoint).map_err(|e| s3_error!(InvalidArgument, "invalid endpoint url: {}", e))?;
match parsed_endpoint.scheme() {
"http" | "https" => {}
other => {
return Err(s3_error!(
InvalidArgument,
"unsupported endpoint scheme: {} (only http and https are allowed)",
other
));
}
}
if let Some(queue_dir) = kv_map.get("queue_dir") {
validate_queue_dir(queue_dir.as_str()).await?;
}
if kv_map.contains_key("client_cert") != kv_map.contains_key("client_key") {
return Err(s3_error!(InvalidArgument, "client_cert and client_key must be specified as a pair"));
}
} else if target_type == AUDIT_MQTT_SUB_SYS {
let endpoint = kv_map
.get(rustfs_config::MQTT_BROKER)
.map(String::as_str)
.ok_or_else(|| s3_error!(InvalidArgument, "broker endpoint is required"))?;
let topic = kv_map
.get(rustfs_config::MQTT_TOPIC)
.map(String::as_str)
.ok_or_else(|| s3_error!(InvalidArgument, "topic is required"))?;
let username = kv_map.get(rustfs_config::MQTT_USERNAME).map(String::as_str);
let password = kv_map.get(rustfs_config::MQTT_PASSWORD).map(String::as_str);
check_mqtt_broker_available(endpoint, topic, username, password)
.await
.map_err(|e| s3_error!(InvalidArgument, "MQTT Broker unavailable: {}", e))?;
if let Some(queue_dir) = kv_map.get("queue_dir") {
validate_queue_dir(queue_dir.as_str()).await?;
if let Some(qos) = kv_map.get("qos") {
match qos.parse::<u8>() {
Ok(1) | Ok(2) => {}
Ok(0) => return Err(s3_error!(InvalidArgument, "qos should be 1 or 2 if queue_dir is set")),
_ => return Err(s3_error!(InvalidArgument, "qos must be an integer 0, 1, or 2")),
}
}
}
}
let mut kvs = rustfs_ecstore::config::KVS::new();
for (key, value) in kv_map {
kvs.insert(key, value);
}
kvs.insert(ENABLE_KEY.to_string(), EnableState::On.to_string());
update_audit_config_and_reload(|config| {
config
.0
.entry(target_type.to_lowercase())
.or_default()
.insert(target_name.to_lowercase(), kvs.clone());
true
})
.await?;
Ok(build_response(StatusCode::OK, Body::empty(), req.headers.get("x-request-id")))
}
}
pub struct ListAuditTargets {}
#[async_trait::async_trait]
impl Operation for ListAuditTargets {
async fn call(&self, req: S3Request<Body>, _params: Params<'_, '_>) -> S3Result<S3Response<(StatusCode, Body)>> {
let span = Span::current();
let _enter = span.enter();
check_permissions(&req).await?;
let mut runtime_statuses = HashMap::new();
if let Some(system) = audit_system() {
let targets = system.get_target_values().await;
let semaphore = Arc::new(Semaphore::new(10));
let mut futures = FuturesUnordered::new();
for target in targets {
let sem = Arc::clone(&semaphore);
futures.push(async move {
let _permit = sem.acquire().await;
let status = match timeout(Duration::from_secs(3), target.is_active()).await {
Ok(Ok(true)) => "online",
_ => "offline",
};
((target.id().id.clone(), target.id().name.to_string()), status.to_string())
});
}
while let Some((key, status)) = futures.next().await {
runtime_statuses.insert(key, status);
}
}
let config = load_server_config_from_store().await?;
let audit_endpoints = merge_audit_endpoints(&config, runtime_statuses);
let data = serde_json::to_vec(&AuditEndpointsResponse { audit_endpoints })
.map_err(|e| s3_error!(InternalError, "failed to serialize audit targets: {}", e))?;
Ok(build_response(StatusCode::OK, Body::from(data), req.headers.get("x-request-id")))
}
}
pub struct RemoveAuditTarget {}
#[async_trait::async_trait]
impl Operation for RemoveAuditTarget {
async fn call(&self, req: S3Request<Body>, params: Params<'_, '_>) -> S3Result<S3Response<(StatusCode, Body)>> {
let span = Span::current();
let _enter = span.enter();
let (target_type, target_name) = extract_target_params(&params)?;
check_permissions(&req).await?;
let config_snapshot = load_server_config_from_store().await?;
if let Some(reason) = audit_target_mutation_block_reason(&config_snapshot, target_type, target_name) {
return Err(s3_error!(InvalidRequest, "{reason}"));
}
update_audit_config_and_reload(|config| {
let mut changed = false;
if let Some(targets) = config.0.get_mut(&target_type.to_lowercase()) {
if targets.remove(&target_name.to_lowercase()).is_some() {
changed = true;
}
if targets.is_empty() {
config.0.remove(&target_type.to_lowercase());
}
}
changed
})
.await?;
Ok(build_response(StatusCode::OK, Body::empty(), req.headers.get("x-request-id")))
}
}
#[cfg(test)]
mod tests {
use super::*;
use matchit::Router;
use rustfs_ecstore::config::{KV, KVS};
use std::collections::{HashMap, HashSet};
use temp_env::{with_var, with_vars, with_vars_unset};
fn enabled_kvs(value: &str) -> KVS {
KVS(vec![KV {
key: ENABLE_KEY.to_string(),
value: value.to_string(),
hidden_if_empty: false,
}])
}
fn with_audit_webhook_target_env_cleared<F>(target_name: &str, f: F)
where
F: FnOnce(),
{
let target_name = target_name.to_ascii_uppercase();
let mut env_keys = vec![format!(
"{ENV_PREFIX}{}{DEFAULT_DELIMITER}{}{DEFAULT_DELIMITER}{target_name}",
AUDIT_WEBHOOK_SUB_SYS.to_ascii_uppercase(),
ENABLE_KEY.to_ascii_uppercase(),
)];
for key in AUDIT_WEBHOOK_KEYS {
let env_key = format!(
"{ENV_PREFIX}{}{DEFAULT_DELIMITER}{}{DEFAULT_DELIMITER}{target_name}",
AUDIT_WEBHOOK_SUB_SYS.to_ascii_uppercase(),
key.to_ascii_uppercase(),
);
if !env_keys.contains(&env_key) {
env_keys.push(env_key);
}
}
with_vars_unset(env_keys, f);
}
#[test]
fn merge_audit_endpoints_marks_config_env_and_mixed_sources() {
let config = Config(HashMap::from([(
AUDIT_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([
("mixed-target".to_string(), enabled_kvs("on")),
("config-target".to_string(), enabled_kvs("on")),
]),
)]));
with_vars(
[
("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_MIXED-TARGET", Some("https://example.com/hook")),
("RUSTFS_AUDIT_WEBHOOK_ENABLE_ENV-ONLY", Some("on")),
("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")),
],
|| {
let runtime = HashMap::from([
(("mixed-target".to_string(), "webhook".to_string()), "online".to_string()),
(("env-only".to_string(), "webhook".to_string()), "online".to_string()),
]);
let merged = merge_audit_endpoints(&config, runtime);
let mixed = merged
.iter()
.find(|entry| entry.account_id == "mixed-target")
.expect("mixed target should be present");
assert_eq!(mixed.source, AuditEndpointSource::Mixed);
let env_only = merged
.iter()
.find(|entry| entry.account_id == "env-only")
.expect("env-only target should be present");
assert_eq!(env_only.source, AuditEndpointSource::Env);
let config_only = merged
.iter()
.find(|entry| entry.account_id == "config-target")
.expect("config target should be present");
assert_eq!(config_only.source, AuditEndpointSource::Config);
},
);
}
#[test]
fn audit_target_mutation_block_reason_rejects_env_managed_target() {
with_vars(
[
("RUSTFS_AUDIT_WEBHOOK_ENABLE_PRIMARY", Some("on")),
("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook")),
],
|| {
let config = Config(HashMap::new());
let reason = audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primary");
assert!(reason.is_some());
assert!(reason.unwrap().contains("managed by environment variables"));
},
);
}
#[test]
fn audit_target_mutation_block_reason_rejects_mixed_target() {
with_var("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook"), || {
let config = Config(HashMap::from([(
AUDIT_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("primary".to_string(), enabled_kvs("on"))]),
)]));
let reason = audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primary");
assert!(reason.is_some());
assert!(reason.unwrap().contains("both persisted config and environment variables"));
});
}
#[test]
fn merge_audit_endpoints_marks_disabled_config_with_env_override_as_mixed() {
let config = Config(HashMap::from([(
AUDIT_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("mixed-disabled".to_string(), enabled_kvs("off"))]),
)]));
with_vars(
[
("RUSTFS_AUDIT_WEBHOOK_ENABLE_MIXED-DISABLED", Some("on")),
("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_MIXED-DISABLED", Some("https://example.com/hook")),
],
|| {
let merged = merge_audit_endpoints(&config, HashMap::new());
let mixed = merged
.iter()
.find(|entry| entry.account_id == "mixed-disabled")
.expect("mixed target should be present");
assert_eq!(mixed.source, AuditEndpointSource::Mixed);
assert_eq!(mixed.status, "offline");
},
);
}
#[test]
fn merge_audit_endpoints_includes_env_only_target_without_runtime_status() {
let config = Config(HashMap::new());
with_vars(
[
("RUSTFS_AUDIT_WEBHOOK_ENABLE_ENV-ONLY", Some("on")),
("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")),
],
|| {
let merged = merge_audit_endpoints(&config, HashMap::new());
let env_only = merged
.iter()
.find(|entry| entry.account_id == "env-only")
.expect("env-only target should be present");
assert_eq!(env_only.source, AuditEndpointSource::Env);
assert_eq!(env_only.status, "offline");
},
);
}
#[test]
fn collect_validated_key_values_rejects_duplicate_keys() {
let allowed_keys: HashSet<&str> = ["endpoint", "auth_token"].into_iter().collect();
let key_values = vec![
KeyValue {
key: "endpoint".to_string(),
value: "https://example.com/one".to_string(),
},
KeyValue {
key: "endpoint".to_string(),
value: "https://example.com/two".to_string(),
},
];
let err = collect_validated_key_values(&key_values, &allowed_keys, AUDIT_WEBHOOK_SUB_SYS).unwrap_err();
assert!(err.to_string().contains("duplicate key"));
}
#[test]
fn collect_validated_key_values_rejects_unsupported_key() {
let allowed_keys: HashSet<&str> = AUDIT_WEBHOOK_KEYS.iter().copied().collect();
let key_values = vec![KeyValue {
key: "not_a_real_key".to_string(),
value: "/tmp/rustfs-audit".to_string(),
}];
let err = collect_validated_key_values(&key_values, &allowed_keys, AUDIT_WEBHOOK_SUB_SYS).unwrap_err();
assert!(err.to_string().contains("not allowed for audit target type"));
}
#[test]
fn extract_target_params_rejects_missing_or_unsupported_values() {
let mut root_router = Router::new();
root_router.insert("/", ()).expect("route should insert");
let missing_type_params = root_router.at("/").expect("route should match");
let missing_type = extract_target_params(&missing_type_params.params).unwrap_err();
assert!(missing_type.to_string().contains("missing required parameter: 'target_type'"));
let mut full_router = Router::new();
full_router
.insert("/v3/audit/target/{target_type}/{target_name}", ())
.expect("route should insert");
let unsupported_type_params = full_router
.at("/v3/audit/target/audit_kafka/primary")
.expect("route should match");
let unsupported_type = extract_target_params(&unsupported_type_params.params).unwrap_err();
assert!(unsupported_type.to_string().contains("unsupported audit target type"));
let mut partial_router = Router::new();
partial_router
.insert("/v3/audit/target/{target_type}", ())
.expect("route should insert");
let missing_name_params = partial_router
.at("/v3/audit/target/audit_webhook")
.expect("route should match");
let missing_name = extract_target_params(&missing_name_params.params).unwrap_err();
assert!(missing_name.to_string().contains("missing required parameter: 'target_name'"));
}
#[test]
fn merge_audit_endpoints_marks_mixed_with_case_insensitive_instance_id() {
let config = Config(HashMap::from([(
AUDIT_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]),
)]));
with_vars(
[
("RUSTFS_AUDIT_WEBHOOK_ENABLE_PRIMARYCASE", Some("on")),
("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_PRIMARYCASE", Some("https://example.com/hook")),
],
|| {
let runtime = HashMap::from([(("PrimaryCase".to_string(), "webhook".to_string()), "online".to_string())]);
let merged = merge_audit_endpoints(&config, runtime);
let mixed = merged
.iter()
.find(|entry| entry.account_id == "PrimaryCase" && entry.service == "webhook")
.expect("mixed target should be present");
assert_eq!(mixed.source, AuditEndpointSource::Mixed);
},
);
}
#[test]
fn audit_target_mutation_block_reason_allows_case_insensitive_config_target_lookup() {
let config = Config(HashMap::from([(
AUDIT_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]),
)]));
with_audit_webhook_target_env_cleared("primarycase", || {
assert!(audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primarycase").is_none());
});
}
#[test]
fn audit_target_mutation_block_reason_allows_runtime_only_target() {
with_audit_webhook_target_env_cleared("primary", || {
let config = Config(HashMap::new());
assert!(audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primary").is_none());
});
}
}
+571 -61
View File
@@ -16,21 +16,23 @@ use crate::admin::router::{AdminOperation, Operation, S3Router};
use crate::auth::{check_key_valid, get_session_token};
use crate::server::ADMIN_PREFIX;
use futures::stream::{FuturesUnordered, StreamExt};
use hashbrown::HashSet as HbHashSet;
use http::{HeaderMap, StatusCode};
use hyper::Method;
use matchit::Params;
use rustfs_config::notify::{NOTIFY_MQTT_SUB_SYS, NOTIFY_WEBHOOK_SUB_SYS};
use rustfs_config::{ENABLE_KEY, EnableState, MAX_ADMIN_REQUEST_BODY_SIZE};
use rustfs_config::notify::{
NOTIFY_MQTT_KEYS, NOTIFY_MQTT_SUB_SYS, NOTIFY_ROUTE_PREFIX, NOTIFY_WEBHOOK_KEYS, NOTIFY_WEBHOOK_SUB_SYS,
};
use rustfs_config::{DEFAULT_DELIMITER, ENABLE_KEY, ENV_PREFIX, EnableState, MAX_ADMIN_REQUEST_BODY_SIZE};
use rustfs_ecstore::config::Config;
use rustfs_targets::check_mqtt_broker_available;
use s3s::{Body, S3Request, S3Response, S3Result, header::CONTENT_TYPE, s3_error};
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
use std::future::Future;
use std::io::{Error, ErrorKind};
use std::net::SocketAddr;
use std::path::Path;
use std::sync::Arc;
use tokio::net::lookup_host;
use tokio::sync::Semaphore;
use tokio::time::{Duration, sleep, timeout};
use tracing::{Span, info, warn};
@@ -80,6 +82,7 @@ struct NotificationEndpoint {
account_id: String,
service: String,
status: String,
source: NotificationEndpointSource,
}
#[derive(Serialize, Debug)]
@@ -87,6 +90,21 @@ struct NotificationEndpointsResponse {
notification_endpoints: Vec<NotificationEndpoint>,
}
type EndpointKey = (String, String);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
enum NotificationEndpointSource {
Config,
Env,
Mixed,
Runtime,
}
fn normalized_endpoint_key(account_id: &str, service: &str) -> EndpointKey {
(account_id.to_lowercase(), service.to_string())
}
// --- Helper Functions ---
async fn check_permissions(req: &S3Request<Body>) -> S3Result<()> {
@@ -155,6 +173,215 @@ async fn validate_queue_dir(queue_dir: &str) -> S3Result<()> {
Ok(())
}
fn config_enable_is_on(value: &str) -> bool {
matches!(value.trim().to_ascii_lowercase().as_str(), "on" | "true" | "yes" | "1")
}
fn collect_configured_endpoint_keys(config: &Config) -> Vec<EndpointKey> {
let mut endpoints = Vec::new();
for (subsystem, service) in [(NOTIFY_WEBHOOK_SUB_SYS, "webhook"), (NOTIFY_MQTT_SUB_SYS, "mqtt")] {
let Some(targets) = config.0.get(subsystem) else {
continue;
};
for (target_name, kvs) in targets {
if target_name == DEFAULT_DELIMITER {
continue;
}
let enabled = kvs.lookup(ENABLE_KEY).as_deref().map(config_enable_is_on).unwrap_or(false);
if enabled {
endpoints.push((target_name.clone(), service.to_string()));
}
}
}
endpoints
}
fn collect_config_entry_keys(config: &Config) -> HbHashSet<EndpointKey> {
let mut endpoints = HbHashSet::new();
for (subsystem, service) in [(NOTIFY_WEBHOOK_SUB_SYS, "webhook"), (NOTIFY_MQTT_SUB_SYS, "mqtt")] {
let Some(targets) = config.0.get(subsystem) else {
continue;
};
for target_name in targets.keys() {
if target_name == DEFAULT_DELIMITER {
continue;
}
endpoints.insert(normalized_endpoint_key(target_name, service));
}
}
endpoints
}
fn collect_env_endpoint_keys() -> HbHashSet<EndpointKey> {
let mut endpoints = HbHashSet::new();
for (service, valid_keys) in [("webhook", NOTIFY_WEBHOOK_KEYS), ("mqtt", NOTIFY_MQTT_KEYS)] {
let env_prefix = format!("{ENV_PREFIX}{NOTIFY_ROUTE_PREFIX}{service}{DEFAULT_DELIMITER}").to_uppercase();
for (key, _value) in std::env::vars() {
let Some(rest) = key.strip_prefix(&env_prefix) else {
continue;
};
let mut parts = rest.rsplitn(2, DEFAULT_DELIMITER);
let instance_id_part = parts.next().unwrap_or(DEFAULT_DELIMITER);
let field_name_part = parts.next();
let (field_name, instance_id) = match field_name_part {
Some(field) => (field.to_lowercase(), instance_id_part.to_lowercase()),
None => (instance_id_part.to_lowercase(), DEFAULT_DELIMITER.to_string()),
};
if instance_id == DEFAULT_DELIMITER || instance_id.is_empty() {
continue;
}
if valid_keys.contains(&field_name.as_str()) {
endpoints.insert(normalized_endpoint_key(&instance_id, service));
}
}
}
endpoints
}
fn classify_notification_endpoint_source(
config_targets: &HbHashSet<EndpointKey>,
env_targets: &HbHashSet<EndpointKey>,
key: &EndpointKey,
) -> NotificationEndpointSource {
match (config_targets.contains(key), env_targets.contains(key)) {
(true, true) => NotificationEndpointSource::Mixed,
(true, false) => NotificationEndpointSource::Config,
(false, true) => NotificationEndpointSource::Env,
(false, false) => NotificationEndpointSource::Runtime,
}
}
fn notification_endpoint_source(config: &Config, target_type: &str, target_name: &str) -> NotificationEndpointSource {
let config_targets = collect_config_entry_keys(config);
let env_targets = collect_env_endpoint_keys();
let service = match target_type {
NOTIFY_WEBHOOK_SUB_SYS => "webhook",
NOTIFY_MQTT_SUB_SYS => "mqtt",
_ => "",
};
let key = normalized_endpoint_key(target_name, service);
classify_notification_endpoint_source(&config_targets, &env_targets, &key)
}
fn target_mutation_block_reason(config: &Config, target_type: &str, target_name: &str) -> Option<String> {
match notification_endpoint_source(config, target_type, target_name) {
NotificationEndpointSource::Env => Some(format!(
"target '{}' is managed by environment variables and cannot be modified from the console",
target_name
)),
NotificationEndpointSource::Mixed => Some(format!(
"target '{}' is configured by both persisted config and environment variables; remove the environment variables first",
target_name
)),
NotificationEndpointSource::Config | NotificationEndpointSource::Runtime => None,
}
}
fn merge_notification_endpoints(config: &Config, runtime_statuses: HashMap<EndpointKey, String>) -> Vec<NotificationEndpoint> {
let mut notification_endpoints = Vec::new();
let mut seen = HashSet::new();
let configured_keys = collect_configured_endpoint_keys(config);
let config_targets = collect_config_entry_keys(config);
let env_targets = collect_env_endpoint_keys();
let mut normalized_runtime_statuses: HashMap<EndpointKey, (String, String, String)> = HashMap::new();
for ((account_id, service), status) in runtime_statuses {
let normalized = normalized_endpoint_key(&account_id, &service);
normalized_runtime_statuses
.entry(normalized)
.or_insert((account_id, service, status));
}
for key in configured_keys {
let normalized = normalized_endpoint_key(&key.0, &key.1);
if !seen.insert(normalized.clone()) {
continue;
}
let status = normalized_runtime_statuses
.remove(&normalized)
.map(|(_, _, status)| status)
.unwrap_or_else(|| "offline".to_string());
let source = classify_notification_endpoint_source(&config_targets, &env_targets, &normalized);
notification_endpoints.push(NotificationEndpoint {
account_id: key.0,
service: key.1,
status,
source,
});
}
for (normalized, (account_id, service, status)) in normalized_runtime_statuses {
if seen.insert(normalized.clone()) {
notification_endpoints.push(NotificationEndpoint {
account_id,
service,
status,
source: classify_notification_endpoint_source(&config_targets, &env_targets, &normalized),
});
}
}
for key in &env_targets {
if !seen.insert(key.clone()) {
continue;
}
notification_endpoints.push(NotificationEndpoint {
account_id: key.0.clone(),
service: key.1.clone(),
status: "offline".to_string(),
source: classify_notification_endpoint_source(&config_targets, &env_targets, key),
});
}
notification_endpoints.sort_by(|a, b| a.service.cmp(&b.service).then_with(|| a.account_id.cmp(&b.account_id)));
notification_endpoints
}
fn collect_online_target_arns(region: &str, target_statuses: Vec<(rustfs_targets::arn::TargetID, String)>) -> Vec<String> {
target_statuses
.into_iter()
.filter_map(|(target_id, status)| (status == "online").then(|| target_id.to_arn(region).to_string()))
.collect()
}
fn collect_validated_key_values(
key_values: &[KeyValue],
allowed_keys: &HashSet<&str>,
target_type: &str,
) -> S3Result<HashMap<String, String>> {
let mut kv_map = HashMap::new();
let mut seen = HashSet::new();
for kv in key_values {
if !allowed_keys.contains(kv.key.as_str()) {
return Err(s3_error!(
InvalidArgument,
"key '{}' not allowed for target type '{}'",
kv.key,
target_type
));
}
if !seen.insert(kv.key.as_str()) {
return Err(s3_error!(InvalidArgument, "duplicate key '{}' in request body", kv.key));
}
kv_map.insert(kv.key.clone(), kv.value.clone());
}
Ok(kv_map)
}
// --- Operations ---
pub struct NotificationTarget {}
@@ -167,6 +394,10 @@ impl Operation for NotificationTarget {
check_permissions(&req).await?;
let ns = get_notification_system()?;
let config_snapshot = ns.config.read().await.clone();
if let Some(reason) = target_mutation_block_reason(&config_snapshot, target_type, target_name) {
return Err(s3_error!(InvalidRequest, "{reason}"));
}
let mut input = req.input;
let body_bytes = input.store_all_limited(MAX_ADMIN_REQUEST_BODY_SIZE).await.map_err(|e| {
@@ -183,37 +414,27 @@ impl Operation for NotificationTarget {
_ => unreachable!(),
};
let kv_map: HashMap<&str, &str> = notification_body
.key_values
.iter()
.map(|kv| (kv.key.as_str(), kv.value.as_str()))
.collect();
// Validate keys
for key in kv_map.keys() {
if !allowed_keys.contains(key) {
return Err(s3_error!(InvalidArgument, "key '{}' not allowed for target type '{}'", key, target_type));
}
}
let kv_map = collect_validated_key_values(&notification_body.key_values, &allowed_keys, target_type)?;
// Type-specific validation
if target_type == NOTIFY_WEBHOOK_SUB_SYS {
let endpoint = kv_map
.get("endpoint")
.map(String::as_str)
.ok_or_else(|| s3_error!(InvalidArgument, "endpoint is required"))?;
let url = Url::parse(endpoint).map_err(|e| s3_error!(InvalidArgument, "invalid endpoint url: {}", e))?;
let host = url
.host_str()
.ok_or_else(|| s3_error!(InvalidArgument, "endpoint missing host"))?;
let port = url
.port_or_known_default()
.ok_or_else(|| s3_error!(InvalidArgument, "endpoint missing port"))?;
let addr = format!("{host}:{port}");
if addr.parse::<SocketAddr>().is_err() && lookup_host(&addr).await.is_err() {
return Err(s3_error!(InvalidArgument, "invalid or unresolvable endpoint address"));
let parsed_endpoint = Url::parse(endpoint).map_err(|e| s3_error!(InvalidArgument, "invalid endpoint url: {}", e))?;
match parsed_endpoint.scheme() {
"http" | "https" => {}
other => {
return Err(s3_error!(
InvalidArgument,
"unsupported endpoint scheme: {} (only http and https are allowed)",
other
));
}
}
if let Some(queue_dir) = kv_map.get("queue_dir") {
validate_queue_dir(queue_dir).await?;
validate_queue_dir(queue_dir.as_str()).await?;
}
if kv_map.contains_key("client_cert") != kv_map.contains_key("client_key") {
return Err(s3_error!(InvalidArgument, "client_cert and client_key must be specified as a pair"));
@@ -221,18 +442,20 @@ impl Operation for NotificationTarget {
} else if target_type == NOTIFY_MQTT_SUB_SYS {
let endpoint = kv_map
.get(rustfs_config::MQTT_BROKER)
.map(String::as_str)
.ok_or_else(|| s3_error!(InvalidArgument, "broker endpoint is required"))?;
let topic = kv_map
.get(rustfs_config::MQTT_TOPIC)
.map(String::as_str)
.ok_or_else(|| s3_error!(InvalidArgument, "topic is required"))?;
let username = kv_map.get(rustfs_config::MQTT_USERNAME).copied();
let password = kv_map.get(rustfs_config::MQTT_PASSWORD).copied();
let username = kv_map.get(rustfs_config::MQTT_USERNAME).map(String::as_str);
let password = kv_map.get(rustfs_config::MQTT_PASSWORD).map(String::as_str);
check_mqtt_broker_available(endpoint, topic, username, password)
.await
.map_err(|e| s3_error!(InvalidArgument, "MQTT Broker unavailable: {}", e))?;
if let Some(queue_dir) = kv_map.get("queue_dir") {
validate_queue_dir(queue_dir).await?;
validate_queue_dir(queue_dir.as_str()).await?;
if let Some(qos) = kv_map.get("qos") {
match qos.parse::<u8>() {
Ok(1) | Ok(2) => {}
@@ -243,24 +466,14 @@ impl Operation for NotificationTarget {
}
}
let mut kvs_vec: Vec<_> = notification_body
.key_values
.into_iter()
.map(|kv| rustfs_ecstore::config::KV {
key: kv.key,
value: kv.value,
hidden_if_empty: false,
})
.collect();
kvs_vec.push(rustfs_ecstore::config::KV {
key: ENABLE_KEY.to_string(),
value: EnableState::On.to_string(),
hidden_if_empty: false,
});
let mut kvs = rustfs_ecstore::config::KVS::new();
for (key, value) in kv_map {
kvs.insert(key, value);
}
kvs.insert(ENABLE_KEY.to_string(), EnableState::On.to_string());
info!("Setting target config for type '{}', name '{}'", target_type, target_name);
ns.set_target_config(target_type, target_name, rustfs_ecstore::config::KVS(kvs_vec))
ns.set_target_config(target_type, target_name, kvs)
.await
.map_err(|e| s3_error!(InternalError, "failed to set target config: {}", e))?;
@@ -278,8 +491,6 @@ impl Operation for ListNotificationTargets {
let ns = get_notification_system()?;
let targets = ns.get_target_values().await;
let target_count = targets.len();
let semaphore = Arc::new(Semaphore::new(10));
let mut futures = FuturesUnordered::new();
@@ -291,18 +502,16 @@ impl Operation for ListNotificationTargets {
Ok(Ok(true)) => "online",
_ => "offline",
};
NotificationEndpoint {
account_id: target.id().id.clone(),
service: target.id().name.to_string(),
status: status.to_string(),
}
((target.id().id.clone(), target.id().name.to_string()), status.to_string())
});
}
let mut notification_endpoints = Vec::with_capacity(target_count);
while let Some(endpoint) = futures.next().await {
notification_endpoints.push(endpoint);
let mut runtime_statuses = HashMap::new();
while let Some((key, status)) = futures.next().await {
runtime_statuses.insert(key, status);
}
let config = ns.config.read().await.clone();
let notification_endpoints = merge_notification_endpoints(&config, runtime_statuses);
let data = serde_json::to_vec(&NotificationEndpointsResponse { notification_endpoints })
.map_err(|e| s3_error!(InternalError, "failed to serialize targets: {}", e))?;
@@ -320,16 +529,32 @@ impl Operation for ListTargetsArns {
check_permissions(&req).await?;
let ns = get_notification_system()?;
let active_targets = ns.get_active_targets().await;
let targets = ns.get_target_values().await;
let region = req
.region
.clone()
.ok_or_else(|| s3_error!(InvalidRequest, "region not found"))?;
let semaphore = Arc::new(Semaphore::new(10));
let mut futures = FuturesUnordered::new();
let data_target_arn_list: Vec<_> = active_targets
.iter()
.map(|id| id.to_arn(region.as_str()).to_string())
.collect();
for target in targets {
let sem = Arc::clone(&semaphore);
futures.push(async move {
let _permit = sem.acquire().await;
let status = match timeout(Duration::from_secs(3), target.is_active()).await {
Ok(Ok(true)) => "online",
_ => "offline",
};
(target.id(), status.to_string())
});
}
let mut target_statuses = Vec::new();
while let Some(target_status) = futures.next().await {
target_statuses.push(target_status);
}
let data_target_arn_list = collect_online_target_arns(region.as_str(), target_statuses);
let data = serde_json::to_vec(&data_target_arn_list)
.map_err(|e| s3_error!(InternalError, "failed to serialize targets: {}", e))?;
@@ -348,6 +573,10 @@ impl Operation for RemoveNotificationTarget {
check_permissions(&req).await?;
let ns = get_notification_system()?;
let config_snapshot = ns.config.read().await.clone();
if let Some(reason) = target_mutation_block_reason(&config_snapshot, target_type, target_name) {
return Err(s3_error!(InvalidRequest, "{reason}"));
}
info!("Removing target config for type '{}', name '{}'", target_type, target_name);
ns.remove_target_config(target_type, target_name)
@@ -372,3 +601,284 @@ fn extract_target_params<'a>(params: &'a Params<'_, '_>) -> S3Result<(&'a str, &
let target_name = extract_param(params, "target_name")?;
Ok((target_type, target_name))
}
#[cfg(test)]
mod tests {
use super::*;
use rustfs_ecstore::config::{KV, KVS};
use rustfs_targets::arn::TargetID;
use std::collections::{HashMap, HashSet};
use temp_env::{with_var, with_vars};
fn enabled_kvs(value: &str) -> KVS {
KVS(vec![KV {
key: ENABLE_KEY.to_string(),
value: value.to_string(),
hidden_if_empty: false,
}])
}
#[test]
fn merge_notification_endpoints_keeps_configured_targets_after_runtime_loss() {
let mut cfg_map = HashMap::new();
cfg_map.insert(
NOTIFY_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("webhook-a".to_string(), enabled_kvs("on"))]),
);
cfg_map.insert(
NOTIFY_MQTT_SUB_SYS.to_string(),
HashMap::from([("mqtt-a".to_string(), enabled_kvs("on"))]),
);
let config = Config(cfg_map);
let runtime = HashMap::from([(("webhook-a".to_string(), "webhook".to_string()), "online".to_string())]);
let merged = merge_notification_endpoints(&config, runtime);
let mqtt = merged
.iter()
.find(|entry| entry.account_id == "mqtt-a" && entry.service == "mqtt")
.expect("mqtt-a should be present");
assert_eq!(mqtt.status, "offline");
assert_eq!(mqtt.source, NotificationEndpointSource::Config);
let webhook = merged
.iter()
.find(|entry| entry.account_id == "webhook-a" && entry.service == "webhook")
.expect("webhook-a should be present");
assert_eq!(webhook.status, "online");
assert_eq!(webhook.source, NotificationEndpointSource::Config);
}
#[test]
fn merge_notification_endpoints_skips_disabled_and_default_entries() {
let mut webhook_targets = HashMap::new();
webhook_targets.insert(DEFAULT_DELIMITER.to_string(), enabled_kvs("on"));
webhook_targets.insert("webhook-disabled".to_string(), enabled_kvs("off"));
webhook_targets.insert("webhook-enabled".to_string(), enabled_kvs("on"));
let config = Config(HashMap::from([(NOTIFY_WEBHOOK_SUB_SYS.to_string(), webhook_targets)]));
let runtime = HashMap::from([
(("webhook-enabled".to_string(), "webhook".to_string()), "online".to_string()),
(("env-only".to_string(), "mqtt".to_string()), "offline".to_string()),
]);
let merged = merge_notification_endpoints(&config, runtime);
let env_only = merged
.iter()
.find(|entry| entry.account_id == "env-only" && entry.service == "mqtt")
.expect("env-only should be present");
assert_eq!(env_only.status, "offline");
assert_eq!(env_only.source, NotificationEndpointSource::Runtime);
let enabled = merged
.iter()
.find(|entry| entry.account_id == "webhook-enabled" && entry.service == "webhook")
.expect("webhook-enabled should be present");
assert_eq!(enabled.status, "online");
assert_eq!(enabled.source, NotificationEndpointSource::Config);
}
#[test]
fn merge_notification_endpoints_marks_env_and_mixed_sources() {
let config = Config(HashMap::from([
(
NOTIFY_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("mixed-target".to_string(), enabled_kvs("on"))]),
),
(
NOTIFY_MQTT_SUB_SYS.to_string(),
HashMap::from([("config-target".to_string(), enabled_kvs("on"))]),
),
]));
with_vars(
[
("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_MIXED-TARGET", Some("https://example.com/hook")),
("RUSTFS_NOTIFY_WEBHOOK_ENABLE_ENV-ONLY", Some("on")),
("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")),
],
|| {
let runtime = HashMap::from([
(("mixed-target".to_string(), "webhook".to_string()), "online".to_string()),
(("env-only".to_string(), "webhook".to_string()), "online".to_string()),
]);
let merged = merge_notification_endpoints(&config, runtime);
let mixed = merged
.iter()
.find(|entry| entry.account_id == "mixed-target")
.expect("mixed target should be present");
assert_eq!(mixed.source, NotificationEndpointSource::Mixed);
let env_only = merged
.iter()
.find(|entry| entry.account_id == "env-only")
.expect("env-only target should be present");
assert_eq!(env_only.source, NotificationEndpointSource::Env);
let config_only = merged
.iter()
.find(|entry| entry.account_id == "config-target")
.expect("config target should be present");
assert_eq!(config_only.source, NotificationEndpointSource::Config);
},
);
}
#[test]
fn target_mutation_block_reason_rejects_env_managed_target() {
with_vars(
[
("RUSTFS_NOTIFY_WEBHOOK_ENABLE_PRIMARY", Some("on")),
("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook")),
],
|| {
let config = Config(HashMap::new());
let reason = target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, "primary");
assert!(reason.is_some());
assert!(reason.unwrap().contains("managed by environment variables"));
},
);
}
#[test]
fn target_mutation_block_reason_rejects_mixed_target() {
with_var("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook"), || {
let config = Config(HashMap::from([(
NOTIFY_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("primary".to_string(), enabled_kvs("on"))]),
)]));
let reason = target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, "primary");
assert!(reason.is_some());
assert!(reason.unwrap().contains("both persisted config and environment variables"));
});
}
#[test]
fn target_mutation_block_reason_allows_config_only_target() {
let target_name = "config-only-target";
let config = Config(HashMap::from([(
NOTIFY_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([(target_name.to_string(), enabled_kvs("on"))]),
)]));
assert!(target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, target_name).is_none());
}
#[test]
fn merge_notification_endpoints_marks_disabled_config_with_env_override_as_mixed() {
let config = Config(HashMap::from([(
NOTIFY_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("mixed-disabled".to_string(), enabled_kvs("off"))]),
)]));
with_vars(
[
("RUSTFS_NOTIFY_WEBHOOK_ENABLE_MIXED-DISABLED", Some("on")),
("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_MIXED-DISABLED", Some("https://example.com/hook")),
],
|| {
let merged = merge_notification_endpoints(&config, HashMap::new());
let mixed = merged
.iter()
.find(|entry| entry.account_id == "mixed-disabled")
.expect("mixed target should be present");
assert_eq!(mixed.source, NotificationEndpointSource::Mixed);
assert_eq!(mixed.status, "offline");
},
);
}
#[test]
fn merge_notification_endpoints_includes_env_only_target_without_runtime_status() {
let config = Config(HashMap::new());
with_vars(
[
("RUSTFS_NOTIFY_WEBHOOK_ENABLE_ENV-ONLY", Some("on")),
("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")),
],
|| {
let merged = merge_notification_endpoints(&config, HashMap::new());
let env_only = merged
.iter()
.find(|entry| entry.account_id == "env-only")
.expect("env-only target should be present");
assert_eq!(env_only.source, NotificationEndpointSource::Env);
assert_eq!(env_only.status, "offline");
},
);
}
#[test]
fn collect_validated_key_values_rejects_duplicate_keys() {
let allowed_keys: HashSet<&str> = ["endpoint", "auth_token"].into_iter().collect();
let key_values = vec![
KeyValue {
key: "endpoint".to_string(),
value: "https://example.com/one".to_string(),
},
KeyValue {
key: "endpoint".to_string(),
value: "https://example.com/two".to_string(),
},
];
let err = collect_validated_key_values(&key_values, &allowed_keys, NOTIFY_WEBHOOK_SUB_SYS).unwrap_err();
assert!(err.to_string().contains("duplicate key"));
}
#[test]
fn merge_notification_endpoints_marks_mixed_with_case_insensitive_instance_id() {
let config = Config(HashMap::from([(
NOTIFY_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]),
)]));
with_vars(
[
("RUSTFS_NOTIFY_WEBHOOK_ENABLE_PRIMARYCASE", Some("on")),
("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARYCASE", Some("https://example.com/hook")),
],
|| {
let runtime = HashMap::from([(("PrimaryCase".to_string(), "webhook".to_string()), "online".to_string())]);
let merged = merge_notification_endpoints(&config, runtime);
let mixed = merged
.iter()
.find(|entry| entry.account_id == "PrimaryCase" && entry.service == "webhook")
.expect("mixed target should be present");
assert_eq!(mixed.source, NotificationEndpointSource::Mixed);
},
);
}
#[test]
fn collect_online_target_arns_filters_offline_targets() {
let arns = collect_online_target_arns(
"us-east-1",
vec![
(TargetID::new("webhook-a".to_string(), "webhook".to_string()), "online".to_string()),
(TargetID::new("mqtt-a".to_string(), "mqtt".to_string()), "offline".to_string()),
],
);
assert_eq!(arns, vec!["arn:rustfs:sqs:us-east-1:webhook-a:webhook".to_string()]);
}
#[test]
fn target_mutation_block_reason_allows_case_insensitive_config_target_lookup() {
let config = Config(HashMap::from([(
NOTIFY_WEBHOOK_SUB_SYS.to_string(),
HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]),
)]));
with_vars(
[
("RUSTFS_NOTIFY_WEBHOOK_ENABLE_PRIMARYCASE", None::<&str>),
("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARYCASE", None::<&str>),
],
|| {
assert!(target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, "primarycase").is_none());
},
);
}
}
+88 -21
View File
@@ -207,30 +207,11 @@ impl Operation for DeleteGroup {
)
.await?;
let group_raw = params
.get("group")
.ok_or_else(|| s3_error!(InvalidArgument, "missing group name in request"))?
.trim();
// Path segments stay percent-encoded in `req.uri.path()` / matchit; IAM uses decoded names (same as GET query).
let group_decoded = percent_decode_str(group_raw)
.decode_utf8()
.map_err(|_| s3_error!(InvalidArgument, "invalid group name encoding"))?;
let group = group_decoded.trim();
// Validate the group name format
if group.is_empty() || group.len() > 256 {
return Err(s3_error!(InvalidArgument, "invalid group name"));
}
// Sanity check the group name
if group.contains(['/', '\\', '\0']) {
return Err(s3_error!(InvalidArgument, "group name contains invalid characters"));
}
let group = decode_delete_group_name(&params)?;
let Ok(iam_store) = rustfs_iam::get() else { return Err(s3_error!(InternalError, "iam not init")) };
let updated_at = iam_store.remove_users_from_group(group, vec![]).await.map_err(|e| {
let updated_at = iam_store.remove_users_from_group(&group, vec![]).await.map_err(|e| {
warn!("delete group failed, e: {:?}", e);
match e {
rustfs_iam::error::Error::GroupNotEmpty => {
@@ -276,6 +257,33 @@ impl Operation for DeleteGroup {
}
}
fn decode_delete_group_name<'a>(params: &'a Params<'_, '_>) -> S3Result<std::borrow::Cow<'a, str>> {
let group_raw = params
.get("group")
.ok_or_else(|| s3_error!(InvalidArgument, "missing group name in request"))?
.trim();
// Path segments stay percent-encoded in `req.uri.path()` / matchit; IAM uses decoded names (same as GET query).
let decoded = percent_decode_str(group_raw)
.decode_utf8()
.map_err(|_| s3_error!(InvalidArgument, "invalid group name encoding"))?;
let group = decoded.trim();
if group.is_empty() || group.len() > 256 {
return Err(s3_error!(InvalidArgument, "invalid group name"));
}
if group.contains(['/', '\\', '\0']) {
return Err(s3_error!(InvalidArgument, "group name contains invalid characters"));
}
if group.len() == decoded.len() {
Ok(decoded)
} else {
Ok(std::borrow::Cow::Owned(group.to_string()))
}
}
pub struct SetGroupStatus {}
#[async_trait::async_trait]
impl Operation for SetGroupStatus {
@@ -484,3 +492,62 @@ impl Operation for UpdateGroupMembers {
Ok(S3Response::with_headers((StatusCode::OK, Body::empty()), header))
}
}
#[cfg(test)]
mod tests {
use super::*;
use matchit::Router;
fn with_delete_group_params<T>(path: &str, f: impl FnOnce(&Params<'_, '_>) -> T) -> T {
let mut router = Router::new();
router
.insert("/rustfs/admin/v3/group/{group}", ())
.expect("route should insert");
let matched = router.at(path).expect("route should match");
f(&matched.params)
}
#[test]
fn decode_delete_group_name_percent_decodes_path_segment() {
let group = with_delete_group_params("/rustfs/admin/v3/group/dev%2Bops%20team", |params| {
decode_delete_group_name(params).map(|group| group.into_owned())
})
.expect("encoded group name should decode");
assert_eq!(group, "dev+ops team");
}
#[test]
fn decode_delete_group_name_rejects_invalid_utf8() {
let err = with_delete_group_params("/rustfs/admin/v3/group/%FF", |params| {
decode_delete_group_name(params).map(|group| group.into_owned())
})
.expect_err("invalid utf-8 should fail");
assert_eq!(err.code(), &S3ErrorCode::InvalidArgument);
assert_eq!(err.message(), Some("invalid group name encoding"));
}
#[test]
fn decode_delete_group_name_rejects_blank_name_after_decoding() {
let err = with_delete_group_params("/rustfs/admin/v3/group/%20", |params| {
decode_delete_group_name(params).map(|group| group.into_owned())
})
.expect_err("blank group should fail");
assert_eq!(err.code(), &S3ErrorCode::InvalidArgument);
assert_eq!(err.message(), Some("invalid group name"));
}
#[test]
fn decode_delete_group_name_rejects_path_separator_after_decoding() {
let err = with_delete_group_params("/rustfs/admin/v3/group/team%2Fops", |params| {
decode_delete_group_name(params).map(|group| group.into_owned())
})
.expect_err("decoded slash should fail");
assert_eq!(err.code(), &S3ErrorCode::InvalidArgument);
assert_eq!(err.message(), Some("group name contains invalid characters"));
}
}
+2
View File
@@ -13,6 +13,7 @@
// limitations under the License.
pub mod account_info;
pub mod audit;
pub mod bucket_meta;
pub mod event;
pub mod group;
@@ -51,6 +52,7 @@ mod tests {
fn test_handler_struct_creation() {
// Test that handler structs can be created
let _account_handler = account_info::AccountInfoHandler {};
let _list_audit_targets = audit::ListAuditTargets {};
let _service_handler = system::ServiceHandle {};
let _server_info_handler = system::ServerInfoHandler {};
let _inspect_data_handler = system::InspectDataHandler {};
+3 -2
View File
@@ -25,8 +25,8 @@ mod console_test;
mod route_registration_test;
use handlers::{
bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts, system,
tier, user,
audit, bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts,
system, tier, user,
};
use router::{AdminOperation, S3Router};
use s3s::route::S3Route;
@@ -55,6 +55,7 @@ pub fn make_admin_route(console_enabled: bool) -> std::io::Result<impl S3Route>
quota::register_quota_route(&mut r)?;
bucket_meta::register_bucket_meta_route(&mut r)?;
audit::register_audit_target_route(&mut r)?;
replication::register_replication_route(&mut r)?;
site_replication::register_site_replication_route(&mut r)?;
+9 -4
View File
@@ -14,8 +14,8 @@
use crate::admin::{
handlers::{
bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts, system,
tier, user,
audit, bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts,
system, tier, user,
},
router::{AdminOperation, S3Router},
};
@@ -50,6 +50,7 @@ fn register_admin_routes(router: &mut S3Router<AdminOperation>) {
tier::register_tier_route(router).expect("register tier route");
quota::register_quota_route(router).expect("register quota route");
bucket_meta::register_bucket_meta_route(router).expect("register bucket meta route");
audit::register_audit_target_route(router).expect("register audit target route");
replication::register_replication_route(router).expect("register replication route");
site_replication::register_site_replication_route(router).expect("register site replication route");
profile_admin::register_profiling_route(router).expect("register profile route");
@@ -60,7 +61,6 @@ fn register_admin_routes(router: &mut S3Router<AdminOperation>) {
#[test]
fn test_register_routes_cover_representative_admin_paths() {
let mut router: S3Router<AdminOperation> = S3Router::new(false);
register_admin_routes(&mut router);
assert_route(&router, Method::GET, HEALTH_PREFIX);
assert_route(&router, Method::HEAD, HEALTH_PREFIX);
@@ -91,6 +91,9 @@ fn test_register_routes_cover_representative_admin_paths() {
assert_route(&router, Method::POST, &admin_path("/v3/idp/builtin/policy/detach"));
assert_route(&router, Method::GET, &admin_path("/v3/idp/builtin/policy-entities"));
assert_route(&router, Method::GET, &admin_path("/v3/target/list"));
assert_route(&router, Method::GET, &admin_path("/v3/audit/target/list"));
assert_route(&router, Method::PUT, &admin_path("/v3/audit/target/audit_webhook/test-audit"));
assert_route(&router, Method::DELETE, &admin_path("/v3/audit/target/audit_webhook/test-audit/reset"));
assert_route(&router, Method::GET, &admin_path("/v3/accountinfo"));
assert_route(&router, Method::POST, &admin_path("/v3/service"));
@@ -165,7 +168,6 @@ fn test_register_routes_cover_representative_admin_paths() {
#[test]
fn test_admin_alias_paths_match_existing_admin_routes() {
let mut router: S3Router<AdminOperation> = S3Router::new(false);
register_admin_routes(&mut router);
for (method, path) in [
@@ -180,6 +182,9 @@ fn test_admin_alias_paths_match_existing_admin_routes() {
(Method::PUT, compat_admin_alias_path("/v3/set-policy")),
(Method::PUT, compat_admin_alias_path("/v3/set-bucket-quota")),
(Method::GET, compat_admin_alias_path("/v3/get-bucket-quota")),
(Method::GET, compat_admin_alias_path("/v3/audit/target/list")),
(Method::PUT, compat_admin_alias_path("/v3/audit/target/audit_webhook/test-audit")),
(Method::DELETE, compat_admin_alias_path("/v3/audit/target/audit_webhook/test-audit/reset")),
(Method::POST, compat_admin_alias_path("/v3/heal/")),
(Method::POST, compat_admin_alias_path("/v3/heal/test-bucket")),
(Method::POST, compat_admin_alias_path("/v3/heal/test-bucket/prefix")),
+2
View File
@@ -1442,6 +1442,7 @@ async fn authorize_replication_extension_request(req: &mut S3Request<Body>, ext_
object: None,
version_id: None,
region: get_global_region(),
..Default::default()
});
license_check().map_err(|er| match er.kind() {
@@ -2163,6 +2164,7 @@ async fn authorize_misc_extension_request(req: &mut S3Request<Body>, route: &Mis
object,
version_id: None,
region: get_global_region(),
..Default::default()
});
license_check().map_err(|er| match er.kind() {
+6 -2
View File
@@ -22,7 +22,7 @@ use crate::auth::get_condition_values;
use crate::error::ApiError;
use crate::server::RemoteAddr;
use crate::storage::access::{ReqInfo, authorize_request, req_info_ref};
use crate::storage::helper::OperationHelper;
use crate::storage::helper::{OperationHelper, spawn_background_with_context};
use crate::storage::s3_api::bucket::{build_list_buckets_output, build_list_objects_v2_output};
use crate::storage::s3_api::common::rustfs_owner;
use crate::storage::s3_api::{acl, encryption, replication, tagging};
@@ -1494,7 +1494,11 @@ impl DefaultBucketUsecase {
&& let Some(store) = new_object_layer_fn()
{
let bucket_name = bucket.clone();
tokio::spawn(async move {
let request_context = req
.extensions
.get::<crate::storage::request_context::RequestContext>()
.cloned();
spawn_background_with_context(request_context, async move {
if let Err(err) = enqueue_transition_for_existing_objects(store, &bucket_name).await {
warn!(bucket = %bucket_name, error = ?err, "failed to enqueue transition for existing objects");
}
@@ -27,8 +27,8 @@ use rustfs_ecstore::{
global::GLOBAL_TierConfigMgr,
store::ECStore,
store_api::{
BucketOperations, BucketOptions, MakeBucketOptions, MultipartOperations, ObjectIO, ObjectOperations, ObjectOptions,
PutObjReader,
BucketOperations, BucketOptions, ChunkNativePutData, MakeBucketOptions, MultipartOperations, ObjectIO, ObjectOperations,
ObjectOptions,
},
tier::{
tier_config::{TierConfig, TierType},
@@ -148,7 +148,7 @@ async fn upload_test_object(
object: &str,
data: &[u8],
) -> rustfs_ecstore::store_api::ObjectInfo {
let mut reader = PutObjReader::from_vec(data.to_vec());
let mut reader = ChunkNativePutData::from_vec(data.to_vec());
(**ecstore)
.put_object(bucket, object, &mut reader, &ObjectOptions::default())
.await
@@ -446,7 +446,7 @@ async fn complete_multipart_upload_transitions_immediately_via_usecase() {
.await
.expect("Failed to create multipart upload");
let mut reader = PutObjReader::from_vec(payload.to_vec());
let mut reader = ChunkNativePutData::from_vec(payload.to_vec());
let uploaded_part = ecstore
.put_object_part(bucket.as_str(), object, &upload.upload_id, 1, &mut reader, &ObjectOptions::default())
.await
+41 -20
View File
@@ -25,6 +25,7 @@ use crate::storage::options::{
copy_src_opts, extract_metadata, get_complete_multipart_upload_opts, get_content_sha256_with_query, get_opts,
parse_copy_source_range, put_opts,
};
use crate::storage::request_context::spawn_traced;
use crate::storage::s3_api::multipart::build_list_parts_output;
use crate::storage::*;
use bytes::Bytes;
@@ -43,9 +44,12 @@ use rustfs_ecstore::compress::is_compressible;
use rustfs_ecstore::error::{StorageError, is_err_object_not_found, is_err_version_not_found};
use rustfs_ecstore::new_object_layer_fn;
use rustfs_ecstore::set_disk::{MAX_PARTS_COUNT, is_valid_storage_class};
use rustfs_ecstore::store_api::{CompletePart, HTTPRangeSpec, MultipartUploadResult, ObjectIO, ObjectOptions, PutObjReader};
use rustfs_ecstore::store_api::{
ChunkNativePutData, CompletePart, HTTPRangeSpec, MultipartUploadResult, ObjectIO, ObjectOptions,
};
use rustfs_ecstore::store_api::{MultipartOperations, ObjectOperations};
use rustfs_filemeta::{ReplicationStatusType, ReplicationType};
use rustfs_object_io::put::PutObjectChecksums;
use rustfs_rio::{CompressReader, HashReader, Reader, WarpReader};
use rustfs_s3_common::S3Operation;
use rustfs_targets::EventName;
@@ -406,7 +410,7 @@ impl DefaultMultipartUsecase {
};
let mpu_version_clone = mpu_version.clone();
let mpu_version_for_event = mpu_version.clone();
tokio::spawn(async move {
spawn_traced(async move {
manager
.invalidate_cache_versioned(&mpu_bucket, &mpu_key, mpu_version_clone.as_deref())
.await;
@@ -718,6 +722,13 @@ impl DefaultMultipartUsecase {
.map_err(ApiError::from)?;
let mut size = size.ok_or_else(|| s3_error!(UnexpectedContent))?;
let mut requested_checksum_type = rustfs_rio::ChecksumType::from_header(&req.headers);
if !requested_checksum_type.is_set()
&& let Some(checksum_algo) = fi.user_defined.get(rustfs_rio::RUSTFS_MULTIPART_CHECKSUM)
&& let Some(checksum_type) = fi.user_defined.get(rustfs_rio::RUSTFS_MULTIPART_CHECKSUM_TYPE)
{
requested_checksum_type = rustfs_rio::ChecksumType::from_string_with_obj_type(checksum_algo, checksum_type);
}
// Apply adaptive buffer sizing based on part size for optimal streaming performance.
// Uses workload profile configuration (enabled by default) to select appropriate buffer size.
@@ -746,11 +757,15 @@ impl DefaultMultipartUsecase {
let mut sha256hex = get_content_sha256_with_query(&req.headers, req.uri.query());
if is_compressible {
let mut hrd = HashReader::new(reader, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?;
let mut hrd =
HashReader::new(reader, size, actual_size, md5hex.take(), sha256hex.take(), false).map_err(ApiError::from)?;
if let Err(err) = hrd.add_checksum_from_s3s(&req.headers, req.trailing_headers.clone(), false) {
return Err(ApiError::from(err).into());
}
if requested_checksum_type.is_set() && hrd.checksum().is_none() {
hrd.enable_auto_checksum(requested_checksum_type).map_err(ApiError::from)?;
}
let compress_reader = CompressReader::new(hrd, CompressionAlgorithm::default());
reader = Box::new(compress_reader);
@@ -764,6 +779,9 @@ impl DefaultMultipartUsecase {
if let Err(err) = reader.add_checksum_from_s3s(&req.headers, req.trailing_headers.clone(), size < 0) {
return Err(ApiError::from(err).into());
}
if requested_checksum_type.is_set() && reader.checksum().is_none() {
reader.enable_auto_checksum(requested_checksum_type).map_err(ApiError::from)?;
}
let has_ssec = sse_customer_algorithm.is_some();
// When SSE-C headers are present, skip managed-encryption metadata to avoid
@@ -823,18 +841,20 @@ impl DefaultMultipartUsecase {
None => (None, None),
};
let mut reader = PutObjReader::new(reader);
let mut reader = ChunkNativePutData::new(reader);
let info = store
.put_object_part(&bucket, &key, &upload_id, part_id, &mut reader, &opts)
.await
.map_err(ApiError::from)?;
let mut checksum_crc32 = input.checksum_crc32;
let mut checksum_crc32c = input.checksum_crc32c;
let mut checksum_sha1 = input.checksum_sha1;
let mut checksum_sha256 = input.checksum_sha256;
let mut checksum_crc64nvme = input.checksum_crc64nvme;
let mut checksums = PutObjectChecksums {
crc32: input.checksum_crc32,
crc32c: input.checksum_crc32c,
sha1: input.checksum_sha1,
sha256: input.checksum_sha256,
crc64nvme: input.checksum_crc64nvme,
};
if let Some(alg) = &input.checksum_algorithm
&& let Some(Some(checksum_str)) = req.trailing_headers.as_ref().map(|trailer| {
@@ -854,25 +874,26 @@ impl DefaultMultipartUsecase {
})
{
match alg.as_str() {
ChecksumAlgorithm::CRC32 => checksum_crc32 = checksum_str,
ChecksumAlgorithm::CRC32C => checksum_crc32c = checksum_str,
ChecksumAlgorithm::SHA1 => checksum_sha1 = checksum_str,
ChecksumAlgorithm::SHA256 => checksum_sha256 = checksum_str,
ChecksumAlgorithm::CRC64NVME => checksum_crc64nvme = checksum_str,
ChecksumAlgorithm::CRC32 => checksums.crc32 = checksum_str,
ChecksumAlgorithm::CRC32C => checksums.crc32c = checksum_str,
ChecksumAlgorithm::SHA1 => checksums.sha1 = checksum_str,
ChecksumAlgorithm::SHA256 => checksums.sha256 = checksum_str,
ChecksumAlgorithm::CRC64NVME => checksums.crc64nvme = checksum_str,
_ => (),
}
}
checksums.merge_from_map(&reader.content_crc());
let output = UploadPartOutput {
server_side_encryption: requested_sse,
ssekms_key_id: requested_kms_key_id,
sse_customer_algorithm,
sse_customer_key_md5,
checksum_crc32,
checksum_crc32c,
checksum_sha1,
checksum_sha256,
checksum_crc64nvme,
checksum_crc32: checksums.crc32,
checksum_crc32c: checksums.crc32c,
checksum_sha1: checksums.sha1,
checksum_sha256: checksums.sha256,
checksum_crc64nvme: checksums.crc64nvme,
e_tag: info.etag.map(|etag| to_s3s_etag(&etag)),
..Default::default()
};
@@ -1190,7 +1211,7 @@ impl DefaultMultipartUsecase {
None => (None, None),
};
let mut reader = PutObjReader::new(reader);
let mut reader = ChunkNativePutData::new(reader);
let dst_opts = ObjectOptions {
user_defined: mp_info.user_defined.clone(),
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,616 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use super::get_object_flow::GetObjectBootstrap;
use super::*;
use crate::app::context::NotifyInterface;
use crate::storage::concurrency::{self, get_buffer_size_opt_in};
use hashbrown::HashMap;
use rustfs_object_io::get::{
CachedGetObjectSource as ObjectIoCachedGetObjectSource, GetObjectBodyPlan as ObjectIoGetObjectBodyPlan,
GetObjectCacheWriteback, GetObjectDataPlaneMetricContract as ObjectIoGetObjectDataPlaneMetricContract, GetObjectFlowResult,
GetObjectResponseMode, MaterializeGetObjectBodyError as ObjectIoMaterializeGetObjectBodyError,
build_cached_get_object_flow_result_from_source as object_io_build_cached_get_object_flow_result_from_source,
finalize_get_object_cache_writeback as object_io_finalize_get_object_cache_writeback,
materialize_get_object_body as object_io_materialize_get_object_body, plan_get_object_body as object_io_plan_get_object_body,
plan_get_object_strategy_layout as object_io_plan_get_object_strategy_layout,
};
pub(super) async fn prepare_get_object_request_context(req: &S3Request<GetObjectInput>) -> S3Result<GetObjectRequestContext> {
let GetObjectInput {
bucket,
key,
version_id,
part_number,
range,
..
} = req.input.clone();
validate_object_key(&key, "GET")?;
let part_number = part_number.map(|v| v as usize);
if let Some(part_num) = part_number
&& part_num == 0
{
return Err(s3_error!(InvalidArgument, "Invalid part number: part number must be greater than 0"));
}
let rs = range.map(|v| match v {
Range::Int { first, last } => HTTPRangeSpec {
is_suffix_length: false,
start: first as i64,
end: if let Some(last) = last { last as i64 } else { -1 },
},
Range::Suffix { length } => HTTPRangeSpec {
is_suffix_length: true,
start: length as i64,
end: -1,
},
});
if rs.is_some() && part_number.is_some() {
return Err(s3_error!(InvalidArgument, "range and part_number invalid"));
}
let opts: ObjectOptions = get_opts(&bucket, &key, version_id.clone(), part_number, &req.headers)
.await
.map_err(ApiError::from)?;
Ok(GetObjectRequestContext {
cache_key: ConcurrencyManager::make_cache_key(&bucket, &key, version_id.as_deref()),
version_id_for_event: version_id.unwrap_or_default(),
bucket,
key,
part_number,
rs,
opts,
headers: req.headers.clone(),
method: req.method.clone(),
sse_customer_key: req.input.sse_customer_key.clone(),
sse_customer_key_md5: req.input.sse_customer_key_md5.clone(),
})
}
impl ObjectIoCachedGetObjectSource for CachedGetObject {
fn body(&self) -> &std::sync::Arc<bytes::Bytes> {
&self.body
}
fn content_length(&self) -> i64 {
self.content_length
}
fn content_type(&self) -> Option<&str> {
self.content_type.as_deref()
}
fn e_tag(&self) -> Option<&str> {
self.e_tag.as_deref()
}
fn last_modified(&self) -> Option<&str> {
self.last_modified.as_deref()
}
fn cache_control(&self) -> Option<&str> {
self.cache_control.as_deref()
}
fn content_disposition(&self) -> Option<&str> {
self.content_disposition.as_deref()
}
fn content_encoding(&self) -> Option<&str> {
self.content_encoding.as_deref()
}
fn content_language(&self) -> Option<&str> {
self.content_language.as_deref()
}
fn storage_class(&self) -> Option<&str> {
self.storage_class.as_deref()
}
fn version_id(&self) -> Option<&str> {
self.version_id.as_deref()
}
fn delete_marker(&self) -> bool {
self.delete_marker
}
fn tag_count(&self) -> Option<i32> {
self.tag_count
}
fn user_metadata(&self) -> &std::collections::HashMap<String, String> {
&self.user_metadata
}
fn checksum_crc32(&self) -> Option<&str> {
self.checksum_crc32.as_deref()
}
fn checksum_crc32c(&self) -> Option<&str> {
self.checksum_crc32c.as_deref()
}
fn checksum_sha1(&self) -> Option<&str> {
self.checksum_sha1.as_deref()
}
fn checksum_sha256(&self) -> Option<&str> {
self.checksum_sha256.as_deref()
}
fn checksum_crc64nvme(&self) -> Option<&str> {
self.checksum_crc64nvme.as_deref()
}
fn checksum_type(&self) -> Option<&ChecksumType> {
self.checksum_type.as_ref()
}
}
pub(super) fn init_get_object_bootstrap(bucket: &str, key: &str, request_id: &str) -> S3Result<GetObjectBootstrap> {
let timeout_config = TimeoutConfig::from_env();
let wrapper = RequestTimeoutWrapper::with_request_id(timeout_config.clone(), request_id.to_string());
let request_start = std::time::Instant::now();
let request_guard = ConcurrencyManager::track_request();
let concurrent_requests = GetObjectGuard::concurrent_requests();
let deadlock_detector = deadlock_detector::get_deadlock_detector();
deadlock_detector.register_request(request_id, format!("GetObject {bucket}/{key}"));
let deadlock_request_guard = DeadlockRequestGuard::new(deadlock_detector, request_id.to_string());
if wrapper.is_timeout() {
warn!(
bucket = %bucket,
key = %key,
timeout_secs = timeout_config.get_object_timeout.as_secs(),
elapsed_ms = wrapper.elapsed().as_millis(),
"GetObject request timed out before processing"
);
return Err(s3_error!(InternalError, "Request timeout before processing"));
}
rustfs_io_metrics::record_get_object_request_start(concurrent_requests);
debug!(
"GetObject request started with {} concurrent requests, timeout={:?}",
concurrent_requests, timeout_config.get_object_timeout
);
Ok(GetObjectBootstrap {
timeout_config,
wrapper,
request_start,
request_guard,
_deadlock_request_guard: deadlock_request_guard,
concurrent_requests,
})
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn maybe_get_cached_get_object_flow_result(
manager: &ConcurrencyManager,
bucket: &str,
key: &str,
cache_key: &str,
version_id_for_event: String,
part_number: Option<usize>,
rs: Option<&HTTPRangeSpec>,
request_start: std::time::Instant,
) -> Option<GetObjectFlowResult> {
if !manager.is_cache_enabled() || part_number.is_some() || rs.is_some() {
return None;
}
let cached = manager.get_cached_object(cache_key).await?;
let cache_serve_duration = request_start.elapsed();
let metric_contract = ObjectIoGetObjectDataPlaneMetricContract::cache_served();
debug!("Serving object from response cache: {} (latency: {:?})", cache_key, cache_serve_duration);
if metric_contract.record_cache_served_metric {
rustfs_io_metrics::record_get_object_cache_served(cache_serve_duration.as_secs_f64(), cached.body.len());
}
rustfs_io_metrics::record_io_path_selected("get", metric_contract.io_path);
rustfs_io_metrics::record_io_copy_mode("get", metric_contract.copy_mode, cached.body.len());
manager.record_transfer(cached.content_length as u64, Duration::from_micros(1));
rustfs_io_metrics::record_get_object(request_start.elapsed().as_millis() as f64, cached.content_length, true);
Some(object_io_build_cached_get_object_flow_result_from_source(
bucket,
key,
cached.as_ref(),
version_id_for_event,
))
}
pub(super) struct GetObjectBodyAdapterOutput {
pub(super) body: Option<StreamingBlob>,
pub(super) body_plan: ObjectIoGetObjectBodyPlan,
pub(super) cache_writeback: Option<GetObjectCacheWriteback>,
}
pub(super) fn spawn_get_object_cache_writeback(
cache_key: &str,
writeback: GetObjectCacheWriteback,
metric_contract: ObjectIoGetObjectDataPlaneMetricContract,
) {
debug_assert_eq!(
metric_contract.request_source,
rustfs_object_io::get::GetObjectDataPlaneRequestSource::Disk
);
debug_assert!(!metric_contract.record_cache_served_metric);
debug_assert!(metric_contract.record_cache_writeback_metric);
let cached_response = CachedGetObject::from_get_object_cache_writeback(writeback);
let cache_key_clone = cache_key.to_string();
crate::storage::request_context::spawn_traced(async move {
let manager = get_concurrency_manager();
manager.put_cached_object(cache_key_clone.clone(), cached_response).await;
debug!("Object cached successfully with metadata: {}", cache_key_clone);
});
if metric_contract.record_cache_writeback_metric {
rustfs_io_metrics::record_object_cache_writeback();
}
}
pub(super) async fn build_get_object_body_adapter<R>(
final_stream: R,
info: &ObjectInfo,
cache_key: &str,
response_content_length: i64,
optimal_buffer_size: usize,
cache_eligibility: rustfs_concurrency::GetObjectCacheEligibility,
) -> S3Result<GetObjectBodyAdapterOutput>
where
R: AsyncRead + Send + Sync + Unpin + 'static,
{
let body_plan = object_io_plan_get_object_body(cache_eligibility, rustfs_config::DEFAULT_OBJECT_SEEK_SUPPORT_THRESHOLD);
match body_plan {
ObjectIoGetObjectBodyPlan::CacheWriteback => {
debug!(
"Reading object into memory for caching: key={} size={}",
cache_key, response_content_length
);
}
ObjectIoGetObjectBodyPlan::BufferSeekable => {
debug!(
"Reading small object into memory for seek support: key={} size={}",
cache_key, response_content_length
);
}
ObjectIoGetObjectBodyPlan::Stream if cache_eligibility.encryption_applied => {
info!(
"Encrypted object: Using unlimited stream for decryption with buffer size {}",
optimal_buffer_size
);
}
_ => {}
}
let materialized =
object_io_materialize_get_object_body(final_stream, info, body_plan, response_content_length, optimal_buffer_size)
.await
.map_err(|err| match err {
ObjectIoMaterializeGetObjectBodyError::CacheRead(err) => {
error!("Failed to read object into memory for caching: {}", err);
ApiError::from(StorageError::other(format!("Failed to read object for caching: {err}")))
}
ObjectIoMaterializeGetObjectBodyError::EncryptedRead(err) => {
error!("Failed to read decrypted object into memory: {}", err);
ApiError::from(StorageError::other(format!("Failed to read decrypted object: {err}")))
}
})?;
Ok(GetObjectBodyAdapterOutput {
body: materialized.body,
body_plan: materialized.plan,
cache_writeback: materialized.cache_writeback.map(|writeback| {
object_io_finalize_get_object_cache_writeback(
info,
writeback,
filter_object_metadata(&info.user_defined).unwrap_or_default(),
)
}),
})
}
pub(super) fn finalize_get_object_completion(
cache_key: &str,
wrapper: &RequestTimeoutWrapper,
timeout_config: &TimeoutConfig,
total_duration: Duration,
response_content_length: i64,
optimal_buffer_size: usize,
metric_contract: ObjectIoGetObjectDataPlaneMetricContract,
) {
rustfs_io_metrics::record_get_object_completion(total_duration.as_secs_f64(), response_content_length, optimal_buffer_size);
rustfs_io_metrics::record_get_object(total_duration.as_millis() as f64, response_content_length, false);
rustfs_io_metrics::record_io_copy_mode("get", metric_contract.copy_mode, response_content_length.max(0) as usize);
if wrapper.is_timeout() {
warn!(
"GetObject request exceeded timeout: key={} duration={:?} timeout={:?}",
cache_key,
wrapper.elapsed(),
timeout_config.get_object_timeout
);
rustfs_io_metrics::record_get_object_timeout(None, Some(wrapper.elapsed().as_secs_f64()));
}
debug!(
"GetObject completed: key={} size={} duration={:?} buffer={}",
cache_key, response_content_length, total_duration, optimal_buffer_size
);
}
#[allow(clippy::too_many_arguments)]
pub(super) fn finalize_get_object_strategy_runtime(
base_buffer_size: usize,
manager: &ConcurrencyManager,
bucket: &str,
key: &str,
info: &ObjectInfo,
rs: Option<&HTTPRangeSpec>,
response_content_length: i64,
permit_wait_duration: Duration,
queue_utilization: f64,
queue_status: &concurrency::IoQueueStatus,
concurrent_requests: usize,
) -> (concurrency::IoStrategy, usize) {
let strategy_layout = object_io_plan_get_object_strategy_layout(
rs,
response_content_length,
0,
get_buffer_size_opt_in(response_content_length),
);
if let Some(range_spec) = rs
&& range_spec.start >= 0
{
manager.record_access(range_spec.start as u64, response_content_length as u64);
}
if response_content_length > 0 {
manager.record_transfer(response_content_length as u64, permit_wait_duration);
}
let io_strategy = manager.calculate_io_strategy_with_context(
info.size,
base_buffer_size,
permit_wait_duration,
strategy_layout.is_sequential_hint,
);
debug!(
wait_ms = permit_wait_duration.as_millis() as u64,
load_level = ?io_strategy.load_level,
buffer_size = io_strategy.buffer_size,
buffer_multiplier = io_strategy.buffer_multiplier,
readahead = io_strategy.enable_readahead,
cache_wb = io_strategy.cache_writeback_enabled,
storage_media = ?io_strategy.storage_media,
access_pattern = ?io_strategy.access_pattern,
bandwidth_tier = ?io_strategy.bandwidth_tier,
concurrent_requests = io_strategy.concurrent_requests,
file_size = info.size,
is_sequential = strategy_layout.is_sequential_hint,
"Enhanced multi-factor I/O strategy calculated"
);
let io_priority = manager.get_io_priority(response_content_length);
if manager.is_priority_scheduling_enabled() {
debug!(
bucket = %bucket,
key = %key,
priority = %io_priority,
request_size = response_content_length,
"I/O priority assigned (based on actual request size)"
);
rustfs_io_metrics::record_io_priority_assignment(io_priority.as_str());
}
rustfs_io_metrics::record_get_object_io_state(
permit_wait_duration.as_secs_f64(),
queue_utilization,
queue_status.permits_in_use,
queue_status.total_permits.saturating_sub(queue_status.permits_in_use),
io_strategy.load_level.as_str(),
io_strategy.buffer_multiplier,
);
rustfs_io_metrics::record_io_priority_assignment(io_priority.as_str());
let strategy_layout = object_io_plan_get_object_strategy_layout(
rs,
response_content_length,
io_strategy.buffer_size,
get_buffer_size_opt_in(response_content_length),
);
debug!(
actual_request_size = response_content_length,
priority = %io_priority.as_str(),
"I/O priority finalized with actual request size"
);
debug!(
"GetObject buffer sizing: file_size={}, base={}, optimal={}, concurrent_requests={}, io_strategy={:?}",
response_content_length,
get_buffer_size_opt_in(response_content_length),
strategy_layout.optimal_buffer_size,
concurrent_requests,
io_strategy.load_level
);
(io_strategy, strategy_layout.optimal_buffer_size)
}
pub(super) fn prepare_put_object_request_context(req: &S3Request<PutObjectInput>) -> PutObjectRequestContext {
PutObjectRequestContext {
headers: req.headers.clone(),
trailing_headers: req.trailing_headers.clone(),
uri_query: req.uri.query().map(str::to_string),
is_post_object: req.extensions.get::<PostObjectRequestMarker>().is_some(),
method: req.method.clone(),
uri: req.uri.clone(),
extensions: req.extensions.clone(),
credentials: req.credentials.clone(),
region: req.region.clone(),
service: req.service.clone(),
}
}
pub(super) fn put_object_execution_context(req: &S3Request<PutObjectInput>) -> (EventName, QuotaOperation, &'static str) {
if req.extensions.get::<PostObjectRequestMarker>().is_some() {
(EventName::ObjectCreatedPost, QuotaOperation::PostObject, "POST")
} else {
(EventName::ObjectCreatedPut, QuotaOperation::PutObject, "PUT")
}
}
pub(super) fn new_operation_helper<T: Send + Sync>(
req: &S3Request<T>,
event_name: EventName,
operation: S3Operation,
suppress_event: bool,
) -> OperationHelper {
let helper = OperationHelper::new(req, event_name, operation);
if suppress_event { helper.suppress_event() } else { helper }
}
pub(super) fn bind_helper_object(
helper: OperationHelper,
object_info: ObjectInfo,
version_id: Option<String>,
) -> OperationHelper {
let helper = helper.object(object_info);
if let Some(version_id) = version_id {
helper.version_id(version_id)
} else {
helper
}
}
pub(super) async fn complete_get_flow_result(
helper: OperationHelper,
request_context: &GetObjectRequestContext,
flow_result: GetObjectFlowResult,
) -> S3Result<S3Response<GetObjectOutput>> {
match flow_result.response_mode {
GetObjectResponseMode::Plain => {
let helper = bind_helper_object(helper, flow_result.event_info, Some(flow_result.version_id_for_event));
let result = Ok(S3Response::new(flow_result.output));
let _ = helper.complete(&result);
result
}
GetObjectResponseMode::CorsWrapped => {
let helper = helper
.object(flow_result.event_info)
.version_id(flow_result.version_id_for_event);
let response = wrap_response_with_cors(
&request_context.bucket,
&request_context.method,
&request_context.headers,
flow_result.output,
)
.await;
let result = Ok(response);
let _ = helper.complete(&result);
result
}
}
}
pub(super) fn complete_put_response(helper: OperationHelper, output: PutObjectOutput) -> S3Result<S3Response<PutObjectOutput>> {
let result = Ok(S3Response::new(output));
let _ = helper.complete(&result);
result
}
#[allow(clippy::too_many_arguments)]
pub(super) fn spawn_put_extract_notification(
notify: Arc<dyn NotifyInterface>,
request_context: Option<crate::storage::request_context::RequestContext>,
bucket: String,
req_params: HashMap<String, String>,
version_id: String,
host: String,
port: u16,
user_agent: String,
obj_info: ObjectInfo,
output: PutObjectOutput,
) {
let event_args = rustfs_notify::EventArgs {
event_name: EventName::ObjectCreatedPut,
bucket_name: bucket,
object: obj_info,
req_params,
resp_elements: extract_resp_elements(&S3Response::new(output)),
version_id,
host,
port,
user_agent,
};
crate::storage::helper::spawn_background_with_context(request_context, async move {
notify.notify(event_args).await;
});
}
pub(super) async fn get_validated_store_adapter(bucket: &str) -> S3Result<Arc<rustfs_ecstore::store::ECStore>> {
get_validated_store(bucket).await
}
pub(super) async fn bucket_prefix_versioning_enabled(bucket: &str, key: &str) -> bool {
BucketVersioningSys::prefix_enabled(bucket, key).await
}
pub(super) async fn authorize_extract_put_target(
request_context: &PutObjectRequestContext,
bucket: &str,
object: &str,
) -> S3Result<()> {
let mut auth_req = S3Request {
input: PutObjectInput::default(),
method: request_context.method.clone(),
uri: request_context.uri.clone(),
headers: request_context.headers.clone(),
extensions: request_context.extensions.clone(),
credentials: request_context.credentials.clone(),
region: request_context.region.clone(),
service: request_context.service.clone(),
trailing_headers: request_context.trailing_headers.clone(),
};
{
let req_info = req_info_mut(&mut auth_req)?;
req_info.bucket = Some(bucket.to_string());
req_info.object = Some(object.to_string());
req_info.version_id = None;
}
authorize_request(&mut auth_req, Action::S3Action(S3Action::PutObjectAction)).await
}
@@ -0,0 +1,280 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use super::DeadlockRequestGuard;
use super::app_adapters::{
bucket_prefix_versioning_enabled, build_get_object_body_adapter, finalize_get_object_completion,
finalize_get_object_strategy_runtime, maybe_get_cached_get_object_flow_result, spawn_get_object_cache_writeback,
};
use super::get_object_zero_copy::{GetObjectPreparedRead, prepare_get_object_read_execution};
use super::types::GetObjectRequestContext;
use crate::error::ApiError;
use crate::storage::concurrency::{self, ConcurrencyManager, GetObjectGuard};
use crate::storage::options::filter_object_metadata;
use crate::storage::timeout_wrapper::{RequestTimeoutWrapper, TimeoutConfig};
use rustfs_ecstore::store_api::{HTTPRangeSpec, ObjectInfo};
use rustfs_object_io::get::{
GetObjectBodyPlan as ObjectIoGetObjectBodyPlan, GetObjectBodySource,
GetObjectDataPlaneMetricContract as ObjectIoGetObjectDataPlaneMetricContract, GetObjectFlowResult, GetObjectOutputContext,
GetObjectReadSetup, build_chunk_blob as object_io_build_chunk_blob,
build_cors_wrapped_get_object_flow_result as object_io_build_cors_wrapped_get_object_flow_result,
build_get_object_checksums as object_io_build_get_object_checksums,
build_get_object_output_context as object_io_build_get_object_output_context,
chunk_body_data_plane_labels as object_io_chunk_body_data_plane_labels,
};
use s3s::S3Result;
use s3s::dto::{ContentType, SSECustomerAlgorithm, SSECustomerKeyMD5, SSEKMSKeyId, ServerSideEncryption, Timestamp};
use std::time::Duration;
pub(super) struct GetObjectBootstrap {
pub(super) timeout_config: TimeoutConfig,
pub(super) wrapper: RequestTimeoutWrapper,
pub(super) request_start: std::time::Instant,
pub(super) request_guard: GetObjectGuard,
pub(super) _deadlock_request_guard: DeadlockRequestGuard,
pub(super) concurrent_requests: usize,
}
#[derive(Clone, Copy)]
pub(super) struct GetObjectFlowRuntime<'a> {
pub(super) manager: &'a ConcurrencyManager,
pub(super) bootstrap: &'a GetObjectBootstrap,
pub(super) base_buffer_size: usize,
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn build_get_object_output_context(
request_context: &GetObjectRequestContext,
cache_key: &str,
manager: &ConcurrencyManager,
bucket: &str,
key: &str,
info: ObjectInfo,
event_info: ObjectInfo,
body_source: GetObjectBodySource,
rs: Option<HTTPRangeSpec>,
content_type: Option<ContentType>,
last_modified: Option<Timestamp>,
response_content_length: i64,
content_range: Option<String>,
server_side_encryption: Option<ServerSideEncryption>,
sse_customer_algorithm: Option<SSECustomerAlgorithm>,
sse_customer_key_md5: Option<SSECustomerKeyMD5>,
ssekms_key_id: Option<SSEKMSKeyId>,
encryption_applied: bool,
permit_wait_duration: Duration,
queue_utilization: f64,
queue_status: &concurrency::IoQueueStatus,
concurrent_requests: usize,
base_buffer_size: usize,
part_number: Option<usize>,
versioned: bool,
) -> S3Result<(GetObjectOutputContext, ObjectIoGetObjectDataPlaneMetricContract)> {
let (io_strategy, optimal_buffer_size) = finalize_get_object_strategy_runtime(
base_buffer_size,
manager,
bucket,
key,
&info,
rs.as_ref(),
response_content_length,
permit_wait_duration,
queue_utilization,
queue_status,
concurrent_requests,
);
let (body, metric_contract) = match body_source {
GetObjectBodySource::Reader(final_stream) => {
let cache_eligibility = manager.get_object_cache_eligibility(
io_strategy.cache_writeback_enabled,
part_number.is_some(),
rs.is_some(),
encryption_applied,
response_content_length,
);
let adapter_output = build_get_object_body_adapter(
final_stream,
&info,
cache_key,
response_content_length,
optimal_buffer_size,
cache_eligibility,
)
.await?;
let metric_contract = ObjectIoGetObjectDataPlaneMetricContract::disk(
rustfs_io_metrics::IoPath::Legacy,
rustfs_io_metrics::CopyMode::SingleCopy,
adapter_output.body_plan,
);
if let Some(writeback) = adapter_output.cache_writeback {
spawn_get_object_cache_writeback(cache_key, writeback, metric_contract);
}
(adapter_output.body, metric_contract)
}
GetObjectBodySource::Chunk {
stream: chunk_stream,
path,
copy_mode,
} => {
let (io_path, copy_mode) = object_io_chunk_body_data_plane_labels(path, copy_mode);
(
object_io_build_chunk_blob(chunk_stream),
ObjectIoGetObjectDataPlaneMetricContract::disk(io_path, copy_mode, ObjectIoGetObjectBodyPlan::Stream),
)
}
};
let checksums = object_io_build_get_object_checksums(&info, &request_context.headers, part_number, rs.as_ref())
.map_err(ApiError::from)?;
let filtered_metadata = filter_object_metadata(&info.user_defined);
Ok((
object_io_build_get_object_output_context(
body,
info,
event_info,
content_type,
last_modified,
response_content_length,
content_range,
server_side_encryption,
sse_customer_algorithm,
sse_customer_key_md5,
ssekms_key_id,
&checksums,
filtered_metadata,
versioned,
optimal_buffer_size,
Some(metric_contract.copy_mode),
),
metric_contract,
))
}
pub(super) async fn run_get_object_flow(
request_context: GetObjectRequestContext,
runtime: GetObjectFlowRuntime<'_>,
) -> S3Result<GetObjectFlowResult> {
let GetObjectFlowRuntime {
manager,
bootstrap,
base_buffer_size,
} = runtime;
let timeout_config = &bootstrap.timeout_config;
let wrapper = &bootstrap.wrapper;
let request_start = bootstrap.request_start;
let concurrent_requests = bootstrap.concurrent_requests;
let bucket = request_context.bucket.clone();
let key = request_context.key.clone();
let cache_key = request_context.cache_key.clone();
let version_id_for_event = request_context.version_id_for_event.clone();
let part_number = request_context.part_number;
let rs = request_context.rs.clone();
let opts = request_context.opts.clone();
if let Some(cached_result) = maybe_get_cached_get_object_flow_result(
manager,
&bucket,
&key,
&cache_key,
version_id_for_event.clone(),
part_number,
rs.as_ref(),
request_start,
)
.await
{
return Ok(cached_result);
}
let prepared_read = prepare_get_object_read_execution(
&request_context,
manager,
wrapper,
timeout_config,
&bucket,
&key,
rs,
&opts,
part_number,
)
.await?;
let GetObjectPreparedRead { io_planning, read_setup } = prepared_read;
let permit_wait_duration = io_planning.permit_wait_duration;
let queue_status = io_planning.queue_status;
let queue_utilization = io_planning.queue_utilization;
let GetObjectReadSetup {
info,
event_info,
body_source,
rs,
content_type,
last_modified,
response_content_length,
content_range,
server_side_encryption,
sse_customer_algorithm,
sse_customer_key_md5,
ssekms_key_id,
encryption_applied,
} = read_setup;
let versioned = bucket_prefix_versioning_enabled(&bucket, &key).await;
let (output_context, metric_contract) = build_get_object_output_context(
&request_context,
&cache_key,
manager,
&bucket,
&key,
info,
event_info,
body_source,
rs,
content_type,
last_modified,
response_content_length,
content_range,
server_side_encryption,
sse_customer_algorithm,
sse_customer_key_md5,
ssekms_key_id,
encryption_applied,
permit_wait_duration,
queue_utilization,
&queue_status,
concurrent_requests,
base_buffer_size,
part_number,
versioned,
)
.await?;
let response_content_length = output_context.response_content_length;
let optimal_buffer_size = output_context.optimal_buffer_size;
let total_duration = request_start.elapsed();
finalize_get_object_completion(
&cache_key,
wrapper,
timeout_config,
total_duration,
response_content_length,
optimal_buffer_size,
metric_contract,
);
Ok(object_io_build_cors_wrapped_get_object_flow_result(output_context, version_id_for_event))
}
@@ -0,0 +1,338 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use super::app_adapters::get_validated_store_adapter;
use super::types::GetObjectRequestContext;
use crate::error::ApiError;
use crate::storage::concurrency::{self, ConcurrencyManager};
use crate::storage::timeout_wrapper::{RequestTimeoutWrapper, TimeoutConfig};
use crate::storage::{
DecryptionRequest, check_preconditions, sse_decryption, validate_sse_headers_for_read, validate_ssec_for_read,
};
use http::HeaderMap;
use rustfs_concurrency::GetObjectQueueSnapshot;
use rustfs_ecstore::store_api::{HTTPRangeSpec, ObjectIO, ObjectOperations, ObjectOptions};
use rustfs_object_io::get::{
ChunkReadDecision, ChunkReadPlanError, GetObjectEncryptionState as ObjectIoGetObjectEncryptionState, GetObjectReadSetup,
build_reader_read_setup as object_io_build_reader_read_setup,
finalize_chunk_read_setup as object_io_finalize_chunk_read_setup,
get_object_chunk_fast_path_guard as object_io_get_object_chunk_fast_path_guard, plan_chunk_read as object_io_plan_chunk_read,
plan_legacy_read as object_io_plan_legacy_read,
};
use rustfs_rio::{Reader, WarpReader};
use s3s::{S3Error, S3ErrorCode, S3Result, s3_error};
use std::time::Duration;
use tracing::{debug, warn};
pub(super) struct GetObjectIoPlanning<'a> {
pub(super) _disk_permit: tokio::sync::SemaphorePermit<'a>,
pub(super) permit_wait_duration: Duration,
pub(super) queue_status: concurrency::IoQueueStatus,
pub(super) queue_utilization: f64,
}
pub(super) struct GetObjectPreparedRead<'a> {
pub(super) io_planning: GetObjectIoPlanning<'a>,
pub(super) read_setup: GetObjectReadSetup,
}
pub(super) async fn acquire_get_object_io_planning<'a>(
manager: &'a ConcurrencyManager,
wrapper: &RequestTimeoutWrapper,
timeout_config: &TimeoutConfig,
bucket: &str,
key: &str,
) -> S3Result<GetObjectIoPlanning<'a>> {
let permit_wait_start = std::time::Instant::now();
let disk_permit = manager
.acquire_disk_read_permit()
.await
.map_err(|_| s3_error!(InternalError, "disk read semaphore closed"))?;
let permit_wait_duration = permit_wait_start.elapsed();
if wrapper.is_timeout() {
warn!(
bucket = %bucket,
key = %key,
wait_ms = permit_wait_duration.as_millis(),
timeout_secs = timeout_config.get_object_timeout.as_secs(),
elapsed_ms = wrapper.elapsed().as_millis(),
"GetObject request timed out while waiting for disk permit"
);
rustfs_io_metrics::record_get_object_timeout(Some("disk_permit"), Some(wrapper.elapsed().as_secs_f64()));
return Err(s3_error!(InternalError, "Request timeout while waiting for disk permit"));
}
let queue_status = manager.io_queue_status();
let queue_snapshot = GetObjectQueueSnapshot::from_available_permits(
queue_status.total_permits,
queue_status.total_permits.saturating_sub(queue_status.permits_in_use),
);
let queue_utilization = queue_snapshot.utilization_percent();
if queue_snapshot.is_congested(80.0) {
warn!(
bucket = %bucket,
key = %key,
queue_utilization = format!("{:.1}%", queue_utilization),
permits_in_use = queue_status.permits_in_use,
total_permits = queue_status.total_permits,
"I/O queue congestion detected"
);
rustfs_io_metrics::record_io_queue_congestion();
}
if wrapper.is_timeout() {
warn!(
bucket = %bucket,
key = %key,
timeout_secs = timeout_config.get_object_timeout.as_secs(),
elapsed_ms = wrapper.elapsed().as_millis(),
"GetObject request timed out before reading object"
);
rustfs_io_metrics::record_get_object_timeout(Some("before_read"), Some(wrapper.elapsed().as_secs_f64()));
return Err(s3_error!(InternalError, "Request timeout before reading object"));
}
Ok(GetObjectIoPlanning {
_disk_permit: disk_permit,
permit_wait_duration,
queue_status,
queue_utilization,
})
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn prepare_get_object_read(
request_context: &GetObjectRequestContext,
store: &rustfs_ecstore::store::ECStore,
manager: &ConcurrencyManager,
bucket: &str,
key: &str,
rs: Option<HTTPRangeSpec>,
h: HeaderMap,
opts: &ObjectOptions,
part_number: Option<usize>,
read_start: std::time::Instant,
) -> S3Result<GetObjectReadSetup> {
let reader = store
.get_object_reader(bucket, key, rs.clone(), h, opts)
.await
.map_err(ApiError::from)?;
let info = reader.object_info;
let read_duration = read_start.elapsed();
rustfs_io_metrics::record_io_path_selected("get", rustfs_io_metrics::IoPath::Legacy);
manager.record_disk_operation(info.size as u64, read_duration, true).await;
check_preconditions(&request_context.headers, &info)?;
debug!(object_size = info.size, part_count = info.parts.len(), "GET object metadata snapshot");
for part in &info.parts {
debug!(
part_number = part.number,
part_size = part.size,
part_actual_size = part.actual_size,
"GET object part details"
);
}
let event_info = info.clone();
validate_sse_headers_for_read(&info.user_defined, &request_context.headers)?;
validate_ssec_for_read(
&info.user_defined,
request_context.sse_customer_key.as_ref(),
request_context.sse_customer_key_md5.as_ref(),
)?;
let read_plan = object_io_plan_legacy_read(&info, rs, part_number).map_err(ApiError::from)?;
debug!(
"GET object metadata check: parts={}, provided_sse_key={:?}",
info.parts.len(),
request_context.sse_customer_key.is_some()
);
let decryption_request = DecryptionRequest {
bucket,
key,
metadata: &info.user_defined,
sse_customer_key: request_context.sse_customer_key.as_ref(),
sse_customer_key_md5: request_context.sse_customer_key_md5.as_ref(),
part_number: None,
parts: &info.parts,
etag: info.etag.as_deref(),
};
let encrypted_stream = reader.stream;
let (encryption_state, final_stream) = match sse_decryption(decryption_request).await? {
Some(material) => {
let server_side_encryption = Some(material.server_side_encryption.clone());
let sse_customer_algorithm = Some(material.algorithm.clone());
let sse_customer_key_md5 = material.customer_key_md5.clone();
let ssekms_key_id = material.kms_key_id.clone();
let (decrypted_stream, plaintext_size) = material
.wrap_reader(encrypted_stream, read_plan.response_content_length)
.await
.map_err(ApiError::from)?;
(
ObjectIoGetObjectEncryptionState {
server_side_encryption,
sse_customer_algorithm,
sse_customer_key_md5,
ssekms_key_id,
encryption_applied: true,
response_content_length_override: Some(plaintext_size),
},
decrypted_stream,
)
}
None => (
ObjectIoGetObjectEncryptionState::default(),
Box::new(WarpReader::new(encrypted_stream)) as Box<dyn Reader>,
),
};
Ok(object_io_build_reader_read_setup(
info,
event_info,
final_stream,
read_plan,
encryption_state,
))
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn prepare_get_object_read_execution<'a>(
request_context: &GetObjectRequestContext,
manager: &'a ConcurrencyManager,
wrapper: &RequestTimeoutWrapper,
timeout_config: &TimeoutConfig,
bucket: &str,
key: &str,
rs: Option<HTTPRangeSpec>,
opts: &ObjectOptions,
part_number: Option<usize>,
) -> S3Result<GetObjectPreparedRead<'a>> {
let h = HeaderMap::new();
let io_planning = acquire_get_object_io_planning(manager, wrapper, timeout_config, bucket, key).await?;
let store = get_validated_store_adapter(bucket).await?;
let read_start = std::time::Instant::now();
let read_setup = match object_io_get_object_chunk_fast_path_guard(
request_context.sse_customer_key.is_some(),
request_context.sse_customer_key_md5.is_some(),
) {
Ok(()) => match prepare_get_object_chunk_read(
request_context,
&store,
manager,
bucket,
key,
rs.clone(),
part_number,
opts,
read_start,
)
.await?
{
Some(read_setup) => read_setup,
None => {
prepare_get_object_read(request_context, &store, manager, bucket, key, rs, h, opts, part_number, read_start)
.await?
}
},
Err(fallback) => {
rustfs_io_metrics::record_io_fallback(fallback.stage, fallback.reason);
prepare_get_object_read(request_context, &store, manager, bucket, key, rs, h, opts, part_number, read_start).await?
}
};
Ok(GetObjectPreparedRead { io_planning, read_setup })
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn prepare_get_object_chunk_read(
request_context: &GetObjectRequestContext,
store: &rustfs_ecstore::store::ECStore,
manager: &ConcurrencyManager,
bucket: &str,
key: &str,
mut rs: Option<HTTPRangeSpec>,
part_number: Option<usize>,
opts: &ObjectOptions,
read_start: std::time::Instant,
) -> S3Result<Option<GetObjectReadSetup>> {
let info = store.get_object_info(bucket, key, opts).await.map_err(ApiError::from)?;
validate_sse_headers_for_read(&info.user_defined, &request_context.headers)?;
validate_ssec_for_read(
&info.user_defined,
request_context.sse_customer_key.as_ref(),
request_context.sse_customer_key_md5.as_ref(),
)?;
check_preconditions(&request_context.headers, &info)?;
let encrypted_object = info.user_defined.contains_key("x-rustfs-encryption-key")
|| info
.user_defined
.contains_key("x-amz-server-side-encryption-customer-algorithm");
if encrypted_object {
rustfs_io_metrics::record_io_fallback(
rustfs_io_metrics::IoStage::ReadSetup,
rustfs_io_metrics::FallbackReason::EncryptionEnabled,
);
return Ok(None);
}
let plan = match object_io_plan_chunk_read(&info, opts.version_id.is_none(), rs.clone(), part_number) {
Ok(ChunkReadDecision::Eligible(plan)) => plan,
Ok(ChunkReadDecision::Fallback(fallback)) => {
rustfs_io_metrics::record_io_fallback(fallback.stage, fallback.reason);
return Ok(None);
}
Err(ChunkReadPlanError::NoSuchKey) => return Err(S3Error::new(S3ErrorCode::NoSuchKey)),
Err(ChunkReadPlanError::MethodNotAllowed) => return Err(S3Error::new(S3ErrorCode::MethodNotAllowed)),
Err(ChunkReadPlanError::Io(err)) => return Err(ApiError::from(err).into()),
};
rs = plan.rs.clone();
let read_duration = read_start.elapsed();
manager.record_disk_operation(info.size as u64, read_duration, true).await;
let event_info = info.clone();
let chunk_result = match store
.get_object_chunks(bucket, key, rs.clone(), HeaderMap::new(), opts)
.await
.map_err(ApiError::from)
{
Ok(result) => result,
Err(_err) => {
rustfs_io_metrics::record_io_fallback(
rustfs_io_metrics::IoStage::HttpBridge,
rustfs_io_metrics::FallbackReason::ChunkBridgeUnavailable,
);
return Ok(None);
}
};
let setup_result = object_io_finalize_chunk_read_setup(info, event_info, chunk_result, plan);
rustfs_io_metrics::record_io_path_selected("get", setup_result.io_path);
Ok(Some(setup_result.read_setup))
}
@@ -0,0 +1,499 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use super::*;
use crate::app::context::NotifyInterface;
use rustfs_object_io::put::{
apply_extract_entry_pax_extensions, apply_trailing_checksums, is_sse_kms_requested, map_extract_archive_error,
normalize_extract_entry_key, resolve_put_object_extract_options,
};
impl DefaultObjectUsecase {
pub(super) async fn run_put_object_extract_flow(
input: PutObjectInput,
request_context: PutObjectRequestContext,
notify: Arc<dyn NotifyInterface>,
resolved_size: i64,
) -> S3Result<PutObjectOutput> {
if is_sse_kms_requested(&input, &request_context.headers) {
return Err(s3_error!(NotImplemented, "SSE-KMS is not supported for extract uploads"));
}
let PutObjectInput {
body,
bucket,
key,
version_id,
cache_control,
content_disposition,
content_encoding,
content_length: _content_length,
content_language,
content_type,
content_md5,
expires,
object_lock_legal_hold_status,
object_lock_mode,
object_lock_retain_until_date,
server_side_encryption,
sse_customer_algorithm,
sse_customer_key,
sse_customer_key_md5,
ssekms_key_id,
storage_class,
tagging,
website_redirect_location,
..
} = input;
let event_version_id = version_id;
let (h_algo, h_key, h_md5) = extract_ssec_params_from_headers(&request_context.headers)?;
let sse_customer_algorithm = sse_customer_algorithm.or(h_algo);
let sse_customer_key = sse_customer_key.or(h_key);
let sse_customer_key_md5 = sse_customer_key_md5.or(h_md5);
let original_sse = server_side_encryption.or(extract_server_side_encryption_from_headers(&request_context.headers)?);
let bucket_sse_config = metadata_sys::get_sse_config(&bucket).await.ok();
let mut effective_sse = original_sse.or_else(|| {
bucket_sse_config.as_ref().and_then(|(config, _timestamp)| {
config.rules.first().and_then(|rule| {
rule.apply_server_side_encryption_by_default
.as_ref()
.map(|sse| match sse.sse_algorithm.as_str() {
"AES256" => ServerSideEncryption::from_static(ServerSideEncryption::AES256),
"aws:kms" => ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS),
_ => ServerSideEncryption::from_static(ServerSideEncryption::AES256),
})
})
})
});
let mut effective_kms_key_id = ssekms_key_id.or_else(|| {
bucket_sse_config.as_ref().and_then(|(config, _timestamp)| {
config.rules.first().and_then(|rule| {
rule.apply_server_side_encryption_by_default
.as_ref()
.and_then(|sse| sse.kms_master_key_id.clone())
})
})
});
if effective_sse
.as_ref()
.is_some_and(|sse| sse.as_str().eq_ignore_ascii_case(ServerSideEncryption::AWS_KMS))
{
return Err(s3_error!(NotImplemented, "SSE-KMS is not supported for extract uploads"));
}
validate_sse_headers_for_write(
effective_sse.as_ref(),
effective_kms_key_id.as_ref(),
sse_customer_algorithm.as_ref(),
sse_customer_key.as_ref(),
sse_customer_key_md5.as_ref(),
true,
)?;
let Some(body) = body else { return Err(s3_error!(IncompleteBody)) };
let size = resolved_size;
validate_object_key(&key, "PUT")?;
let buffer_size = get_buffer_size_opt_in(size);
let body = tokio::io::BufReader::with_capacity(
buffer_size,
StreamReader::new(body.map(|f| f.map_err(|e| std::io::Error::other(e.to_string())))),
);
let Some(ext) = Path::new(&key).extension().and_then(|s| s.to_str()) else {
return Err(s3_error!(InvalidArgument, "key extension not found"));
};
let ext = ext.to_owned();
let md5hex = if let Some(base64_md5) = content_md5 {
let md5 = base64_simd::STANDARD
.decode_to_vec(base64_md5.as_bytes())
.map_err(|e| ApiError::from(StorageError::other(format!("Invalid content MD5: {e}"))))?;
Some(hex_simd::encode_to_string(&md5, hex_simd::AsciiCase::Lower))
} else {
None
};
let sha256hex = get_content_sha256_with_query(&request_context.headers, request_context.uri_query.as_deref());
let actual_size = size;
let mut archive_reader =
HashReader::from_stream(body, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?;
if let Err(err) =
archive_reader.add_checksum_from_s3s(&request_context.headers, request_context.trailing_headers.clone(), false)
{
return Err(ApiError::from(err).into());
}
let archive_etag = Arc::new(Mutex::new(None));
let decoder = CompressionFormat::from_extension(&ext)
.get_decoder(ExtractArchiveEtagReader::new(archive_reader, archive_etag.clone()))
.map_err(|e| {
error!("get_decoder err {:?}", e);
s3_error!(InvalidArgument, "get_decoder err")
})?;
let mut ar = Archive::new(decoder);
let mut entries = ar.entries().map_err(|e| {
error!("get entries err {:?}", e);
s3_error!(InvalidArgument, "get entries err")
})?;
let Some(store) = new_object_layer_fn() else {
return Err(S3Error::with_message(S3ErrorCode::InternalError, "Not init".to_string()));
};
let extract_options = resolve_put_object_extract_options(&request_context.headers);
let version_id = match event_version_id {
Some(v) => v.to_string(),
None => String::new(),
};
let req_params = extract_params_header(&request_context.headers);
let host = get_request_host(&request_context.headers);
let port = get_request_port(&request_context.headers);
let user_agent = get_request_user_agent(&request_context.headers);
let tracing_context = request_context
.extensions
.get::<crate::storage::request_context::RequestContext>()
.cloned();
while let Some(entry) = entries.next().await {
let mut f = match entry {
Ok(f) => f,
Err(e) => {
if extract_options.ignore_errors {
warn!("Skipping archive entry because read failed and ignore-errors is enabled: {e}");
continue;
}
error!("Failed to read archive entry: {}", e);
return Err(s3_error!(InvalidArgument, "Failed to read archive entry: {:?}", e));
}
};
let fpath = match f.path() {
Ok(path) => path,
Err(e) => {
if extract_options.ignore_errors {
warn!("Skipping archive entry because path decode failed and ignore-errors is enabled: {e}");
continue;
}
return Err(s3_error!(InvalidArgument, "Failed to decode archive entry path"));
}
};
let is_dir = f.header().entry_type().is_dir();
let fpath = normalize_extract_entry_key(&fpath.to_string_lossy(), extract_options.prefix.as_deref(), is_dir);
authorize_extract_put_target(&request_context, &bucket, &fpath).await?;
let mut size = f.header().size().unwrap_or_default() as i64;
let archive_entry_mod_time = f
.header()
.mtime()
.ok()
.and_then(|modified_at_secs| OffsetDateTime::from_unix_timestamp(modified_at_secs as i64).ok());
let mut metadata = HashMap::new();
apply_put_request_metadata(
&mut metadata,
&request_context.headers,
&fpath,
cache_control.clone(),
content_disposition.clone(),
content_encoding.clone(),
content_language.clone(),
content_type.clone(),
expires.clone(),
website_redirect_location.clone(),
tagging.clone(),
storage_class.clone(),
)?;
let mut opts = put_opts(&bucket, &fpath, None, &request_context.headers, metadata.clone())
.await
.map_err(ApiError::from)?;
apply_extract_entry_pax_extensions(&mut f, &mut metadata, &mut opts).await?;
if archive_entry_mod_time.is_some() {
opts.mod_time = archive_entry_mod_time;
}
debug!("Extracting file: {}, size: {} bytes", fpath, size);
if is_dir {
if extract_options.ignore_dirs {
debug!("Skipping directory entry during archive extract: {}", fpath);
continue;
}
size = 0;
}
let actual_size = size;
let should_compress = !is_dir && is_compressible(&HeaderMap::new(), &fpath) && size > MIN_COMPRESSIBLE_SIZE as i64;
let mut hrd = if is_dir {
HashReader::from_stream(std::io::Cursor::new(Vec::new()), size, actual_size, None, None, false)
.map_err(ApiError::from)?
} else if should_compress {
insert_str(&mut metadata, SUFFIX_COMPRESSION, CompressionAlgorithm::default().to_string());
insert_str(&mut metadata, SUFFIX_ACTUAL_SIZE, size.to_string());
let hrd = HashReader::from_stream(f, size, actual_size, None, None, false).map_err(ApiError::from)?;
size = HashReader::SIZE_PRESERVE_LAYER;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
size,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(f, size, actual_size, None, None, false).map_err(ApiError::from)?
};
apply_put_request_object_lock_opts(
&bucket,
object_lock_legal_hold_status.clone(),
object_lock_mode.clone(),
object_lock_retain_until_date.clone(),
&mut opts,
)
.await?;
if let Some(material) = sse_encryption(EncryptionRequest {
bucket: &bucket,
key: &fpath,
server_side_encryption: effective_sse.clone(),
ssekms_key_id: effective_kms_key_id.clone(),
sse_customer_algorithm: sse_customer_algorithm.clone(),
sse_customer_key: sse_customer_key.clone(),
sse_customer_key_md5: sse_customer_key_md5.clone(),
content_size: actual_size,
part_number: None,
part_key: None,
part_nonce: None,
})
.await?
{
effective_sse = Some(material.server_side_encryption.clone());
effective_kms_key_id = material.kms_key_id.clone();
let encrypted_reader = material.wrap_reader(hrd);
hrd = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
let encryption_metadata = material.metadata;
metadata.extend(encryption_metadata.clone());
opts.user_defined.extend(encryption_metadata);
}
opts.user_defined.extend(metadata);
let mut reader = rustfs_ecstore::store_api::ChunkNativePutData::new(hrd);
let obj_info = match store.put_object(&bucket, &fpath, &mut reader, &opts).await {
Ok(info) => info,
Err(e) => {
if extract_options.ignore_errors {
warn!("Skipping archive entry because object write failed and ignore-errors is enabled: {e}");
continue;
}
return Err(ApiError::from(e).into());
}
};
let manager = get_concurrency_manager();
let fpath_clone = fpath.clone();
let bucket_clone = bucket.clone();
crate::storage::request_context::spawn_traced(async move {
manager.invalidate_cache_versioned(&bucket_clone, &fpath_clone, None).await;
});
let e_tag = obj_info.etag.clone().map(|etag| to_s3s_etag(&etag));
let output = PutObjectOutput {
e_tag,
..Default::default()
};
spawn_put_extract_notification(
notify.clone(),
tracing_context.clone(),
bucket.clone(),
req_params.clone(),
version_id.clone(),
host.clone(),
port,
user_agent.clone(),
obj_info.clone(),
output,
);
}
let mut checksums = PutObjectChecksums {
crc32: input.checksum_crc32,
crc32c: input.checksum_crc32c,
sha1: input.checksum_sha1,
sha256: input.checksum_sha256,
crc64nvme: input.checksum_crc64nvme,
};
apply_trailing_checksums(
input.checksum_algorithm.as_ref().map(|a| a.as_str()),
&request_context.trailing_headers,
&mut checksums,
);
drop(entries);
let mut decoder = match ar.into_inner() {
Ok(decoder) => decoder,
Err(_) => return Err(s3_error!(InvalidArgument, "Failed to finalize archive reader")),
};
tokio::io::copy(&mut decoder, &mut tokio::io::sink())
.await
.map_err(map_extract_archive_error)?;
let archive_etag = archive_etag
.lock()
.ok()
.and_then(|etag| etag.clone())
.map(|etag| to_s3s_etag(&etag));
let output = PutObjectOutput {
e_tag: archive_etag,
checksum_crc32: checksums.crc32,
checksum_crc32c: checksums.crc32c,
checksum_sha1: checksums.sha1,
checksum_sha256: checksums.sha256,
checksum_crc64nvme: checksums.crc64nvme,
..Default::default()
};
Ok(output)
}
}
#[cfg(test)]
mod tests {
use super::*;
use http::{Extensions, HeaderMap, HeaderValue, Method, Uri};
use rustfs_utils::http::headers::{AMZ_SERVER_SIDE_ENCRYPTION, AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, AMZ_SNOWBALL_EXTRACT};
fn build_request<T>(input: T, method: Method) -> S3Request<T> {
S3Request {
input,
method,
uri: Uri::from_static("/"),
headers: HeaderMap::new(),
extensions: Extensions::new(),
credentials: None,
region: None,
service: None,
trailing_headers: None,
}
}
#[tokio::test]
async fn execute_put_object_rejects_post_object_sse_kms_from_input() {
let input = PutObjectInput::builder()
.bucket("test-bucket".to_string())
.key("test-key".to_string())
.server_side_encryption(Some(ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS)))
.build()
.unwrap();
let mut req = build_request(input, Method::POST);
req.extensions.insert(PostObjectRequestMarker);
let usecase = DefaultObjectUsecase::without_context();
let fs = FS::new();
let err = usecase.execute_put_object(&fs, req).await.unwrap_err();
assert_eq!(err.code(), &S3ErrorCode::NotImplemented);
}
#[tokio::test]
async fn execute_put_object_rejects_extract_sse_kms() {
let input = PutObjectInput::builder()
.bucket("test-bucket".to_string())
.key("archive.tar".to_string())
.server_side_encryption(Some(ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS)))
.build()
.unwrap();
let mut req = build_request(input, Method::PUT);
req.headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true"));
let usecase = DefaultObjectUsecase::without_context();
let fs = FS::new();
let err = usecase.execute_put_object(&fs, req).await.unwrap_err();
assert_eq!(err.code(), &S3ErrorCode::NotImplemented);
}
#[tokio::test]
async fn execute_put_object_extract_rejects_invalid_storage_class() {
let input = PutObjectInput::builder()
.bucket("test-bucket".to_string())
.key("archive.tar".to_string())
.storage_class(Some(StorageClass::from_static("INVALID")))
.build()
.unwrap();
let mut req = build_request(input, Method::PUT);
req.headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true"));
let usecase = DefaultObjectUsecase::without_context();
let fs = FS::new();
let err = usecase.execute_put_object(&fs, req).await.unwrap_err();
assert_eq!(err.code(), &S3ErrorCode::InvalidStorageClass);
}
#[tokio::test]
async fn execute_put_object_rejects_post_object_sse_kms_from_headers() {
let input = PutObjectInput::builder()
.bucket("test-bucket".to_string())
.key("test-key".to_string())
.build()
.unwrap();
let mut req = build_request(input, Method::POST);
req.extensions.insert(PostObjectRequestMarker);
req.headers
.insert(AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static("aws:kms"));
let usecase = DefaultObjectUsecase::without_context();
let fs = FS::new();
let err = usecase.execute_put_object(&fs, req).await.unwrap_err();
assert_eq!(err.code(), &S3ErrorCode::NotImplemented);
}
#[tokio::test]
async fn execute_put_object_rejects_post_object_sse_kms_key_id_header() {
let input = PutObjectInput::builder()
.bucket("test-bucket".to_string())
.key("test-key".to_string())
.build()
.unwrap();
let mut req = build_request(input, Method::POST);
req.extensions.insert(PostObjectRequestMarker);
req.headers
.insert(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, HeaderValue::from_static("test-kms-key-id"));
let usecase = DefaultObjectUsecase::without_context();
let fs = FS::new();
let err = usecase.execute_put_object(&fs, req).await.unwrap_err();
assert_eq!(err.code(), &S3ErrorCode::NotImplemented);
}
}
@@ -0,0 +1,868 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use super::*;
use bytes::Buf;
use futures::{Stream, StreamExt};
use rustfs_ecstore::config::GLOBAL_STORAGE_CLASS;
use rustfs_io_core::{BytesPool, PooledBuffer};
use rustfs_object_io::put::{
PutObjectChecksums, PutObjectIngressKind, PutObjectLegacyHashStagePlan, PutObjectLegacyHashValues, PutObjectTransformStage,
apply_trailing_checksums, build_put_object_ingress_source, build_put_object_legacy_hash_stage,
build_put_object_plain_hash_stage, plan_put_object_body_with_transforms, resolve_put_transformed_fallback_reason,
};
use rustfs_rio::{BlockReadable, BoxReadBlockFuture, EtagResolvable, HashReaderDetector, TryGetIndex};
use rustfs_utils::http::headers::AMZ_TRAILER;
const DEFAULT_SMALL_PUT_EAGER_MAX_BYTES: i64 = 1024 * 1024;
const ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES: &str = "RUSTFS_PUT_SMALL_EAGER_MAX_BYTES";
const ENV_RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER: &str = "RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER";
const SLOW_PUT_PHASE_DEBUG_THRESHOLD_MS: u64 = 100;
const SLOW_PUT_PHASE_WARN_THRESHOLD_MS: u64 = 1_000;
const SLOW_PUT_PHASE_ERROR_THRESHOLD_MS: u64 = 5_000;
fn resolved_checksum_bytes(checksums: &PutObjectChecksums) -> Option<bytes::Bytes> {
[
(rustfs_rio::ChecksumType::CRC32, checksums.crc32.as_deref()),
(rustfs_rio::ChecksumType::CRC32C, checksums.crc32c.as_deref()),
(rustfs_rio::ChecksumType::SHA1, checksums.sha1.as_deref()),
(rustfs_rio::ChecksumType::SHA256, checksums.sha256.as_deref()),
(rustfs_rio::ChecksumType::CRC64_NVME, checksums.crc64nvme.as_deref()),
]
.into_iter()
.find_map(|(checksum_type, value)| {
value.and_then(|value| rustfs_rio::Checksum::new_with_type(checksum_type, value).map(|checksum| checksum.to_bytes(&[])))
})
}
fn clamp_small_put_eager_max_bytes(inline_object_limit_bytes: Option<usize>) -> i64 {
inline_object_limit_bytes
.unwrap_or(DEFAULT_SMALL_PUT_EAGER_MAX_BYTES as usize)
.min(DEFAULT_SMALL_PUT_EAGER_MAX_BYTES as usize) as i64
}
fn env_flag_enabled(name: &str) -> bool {
rustfs_utils::get_env_bool(name, false)
}
fn env_non_negative_i64(name: &str) -> Option<i64> {
rustfs_utils::get_env_opt_i64(name).filter(|value| *value >= 0)
}
fn topology_aware_small_put_eager_max_bytes(store: &rustfs_ecstore::store::ECStore, versioned: bool) -> i64 {
let Some(first_pool) = store.pools.first() else {
return DEFAULT_SMALL_PUT_EAGER_MAX_BYTES;
};
let data_shards = first_pool
.set_drive_count
.saturating_sub(first_pool.default_parity_count)
.max(1);
let inline_object_limit = GLOBAL_STORAGE_CLASS
.get()
.map(|config| config.inline_object_limit_bytes(data_shards, versioned));
clamp_small_put_eager_max_bytes(inline_object_limit)
}
fn resolved_small_put_eager_max_bytes(default_max_bytes: i64) -> i64 {
if env_flag_enabled(ENV_RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER) {
return 0;
}
env_non_negative_i64(ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES)
.map(|value| value.min(DEFAULT_SMALL_PUT_EAGER_MAX_BYTES).min(default_max_bytes))
.unwrap_or(default_max_bytes)
}
fn should_use_small_put_eager_path(size: i64, eager_max_bytes: i64, compression_enabled: bool, encryption_enabled: bool) -> bool {
size > 0 && size <= eager_max_bytes && !compression_enabled && !encryption_enabled
}
fn request_uses_trailing_checksum(headers: &HeaderMap, trailing_headers: &Option<s3s::TrailingHeaders>) -> bool {
trailing_headers.is_some()
|| headers.contains_key(AMZ_TRAILER)
|| matches!(
rustfs_rio::get_content_checksum(headers),
Ok(Some(checksum)) if checksum.checksum_type.trailing()
)
}
fn put_path_label(small_eager: bool, reduced_copy: bool, compressed: bool) -> &'static str {
if small_eager {
"small_eager"
} else if compressed {
"compressed"
} else if reduced_copy {
"reduced_copy"
} else {
"legacy_plain"
}
}
#[allow(clippy::too_many_arguments)]
fn log_put_flow_phase(
bucket: &str,
key: &str,
phase: &str,
elapsed: std::time::Duration,
object_size: i64,
small_eager: bool,
reduced_copy: bool,
compressed: bool,
encrypted: bool,
) {
let duration_ms = elapsed.as_millis() as u64;
if duration_ms < SLOW_PUT_PHASE_DEBUG_THRESHOLD_MS {
return;
}
let put_path = put_path_label(small_eager, reduced_copy, compressed);
if duration_ms >= SLOW_PUT_PHASE_ERROR_THRESHOLD_MS {
error!(
phase,
duration_ms, object_size, put_path, compressed, encrypted, bucket, key, "Small PUT phase is critically slow"
);
} else if duration_ms >= SLOW_PUT_PHASE_WARN_THRESHOLD_MS {
warn!(
phase,
duration_ms, object_size, put_path, compressed, encrypted, bucket, key, "Small PUT phase is slow"
);
} else {
debug!(
phase,
duration_ms, object_size, put_path, compressed, encrypted, bucket, key, "Small PUT phase exceeded debug threshold"
);
}
}
struct PooledBufferReader {
buffer: PooledBuffer,
position: usize,
}
impl PooledBufferReader {
fn new(buffer: PooledBuffer) -> Self {
Self { buffer, position: 0 }
}
}
impl tokio::io::AsyncRead for PooledBufferReader {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let remaining = &self.buffer[self.position..];
if remaining.is_empty() {
return std::task::Poll::Ready(Ok(()));
}
let to_copy = remaining.len().min(buf.remaining());
buf.put_slice(&remaining[..to_copy]);
self.position += to_copy;
std::task::Poll::Ready(Ok(()))
}
}
impl BlockReadable for PooledBufferReader {
fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> {
Box::pin(async move {
let remaining = &self.buffer[self.position..];
if remaining.is_empty() {
return Ok(0);
}
let to_copy = remaining.len().min(buf.len());
buf[..to_copy].copy_from_slice(&remaining[..to_copy]);
self.position += to_copy;
Ok(to_copy)
})
}
}
impl EtagResolvable for PooledBufferReader {}
impl HashReaderDetector for PooledBufferReader {}
impl TryGetIndex for PooledBufferReader {}
async fn read_small_put_body_eager<S, B, E>(body: S, size: i64, pool: std::sync::Arc<BytesPool>) -> S3Result<PooledBuffer>
where
S: Stream<Item = Result<B, E>>,
B: Buf,
E: std::fmt::Display,
{
let expected_len = usize::try_from(size).map_err(|_| s3_error!(InvalidRequest, "Object size overflow"))?;
let mut data = pool.acquire_buffer(expected_len).await;
let mut body = Box::pin(body);
while let Some(result) = body.next().await {
let mut chunk = result.map_err(|err| S3Error::with_message(S3ErrorCode::IncompleteBody, err.to_string()))?;
let chunk_len = chunk.remaining();
if chunk_len == 0 {
continue;
}
let new_len = data
.len()
.checked_add(chunk_len)
.ok_or_else(|| s3_error!(InvalidRequest, "Object size overflow"))?;
if new_len > expected_len {
return Err(s3_error!(IncompleteBody));
}
let start = data.len();
data.resize(new_len, 0);
chunk.copy_to_slice(&mut data[start..new_len]);
if data.len() == expected_len {
return Ok(data);
}
}
if data.len() != expected_len {
return Err(s3_error!(IncompleteBody));
}
Ok(data)
}
async fn build_small_put_eager_hash_stage<S, B, E>(
body: S,
size: i64,
pool: std::sync::Arc<BytesPool>,
hash_values: PutObjectLegacyHashValues,
headers: &HeaderMap,
trailing_headers: Option<s3s::TrailingHeaders>,
) -> S3Result<rustfs_object_io::put::PutObjectHashStage>
where
S: Stream<Item = Result<B, E>>,
B: Buf,
E: std::fmt::Display,
{
let data = read_small_put_body_eager(body, size, pool).await?;
build_put_object_legacy_hash_stage(
Box::new(PooledBufferReader::new(data)),
hash_values,
PutObjectLegacyHashStagePlan {
size,
actual_size: size,
apply_s3_checksum: true,
ignore_s3_checksum_value: false,
},
headers,
trailing_headers,
)
.map_err(ApiError::from)
.map_err(Into::into)
}
impl DefaultObjectUsecase {
pub(super) async fn run_put_object_flow(
input: PutObjectInput,
request_context: PutObjectRequestContext,
request_method_name: &'static str,
resolved_size: i64,
) -> S3Result<PutObjectFlowResult> {
let start_time = std::time::Instant::now();
let PutObjectInput {
body,
bucket,
cache_control,
key,
content_length: _content_length,
content_disposition,
content_encoding,
content_language,
content_type,
expires,
tagging,
metadata,
version_id,
server_side_encryption,
sse_customer_algorithm,
sse_customer_key,
sse_customer_key_md5,
ssekms_key_id,
content_md5,
object_lock_legal_hold_status,
object_lock_mode,
object_lock_retain_until_date,
storage_class,
website_redirect_location,
..
} = input;
let (h_algo, h_key, h_md5) = extract_ssec_params_from_headers(&request_context.headers)?;
let sse_customer_algorithm = sse_customer_algorithm.or(h_algo);
let sse_customer_key = sse_customer_key.or(h_key);
let sse_customer_key_md5 = sse_customer_key_md5.or(h_md5);
let server_side_encryption =
server_side_encryption.or(extract_server_side_encryption_from_headers(&request_context.headers)?);
validate_object_key(&key, request_method_name)?;
let Some(body) = body else { return Err(s3_error!(IncompleteBody)) };
let mut size = resolved_size;
let mut transform_stage = PutObjectTransformStage::default();
let mut plain_reduced_copy_stage = false;
let mut small_object_eager_stage = false;
let bytes_pool = get_concurrency_manager().bytes_pool();
let store = get_validated_store_adapter(&bucket).await?;
let bucket_sse_config = metadata_sys::get_sse_config(&bucket).await.ok();
let mut effective_sse = server_side_encryption.or_else(|| {
bucket_sse_config.as_ref().and_then(|(config, _timestamp)| {
config.rules.first().and_then(|rule| {
rule.apply_server_side_encryption_by_default
.as_ref()
.map(|sse| match sse.sse_algorithm.as_str() {
"AES256" => ServerSideEncryption::from_static(ServerSideEncryption::AES256),
"aws:kms" => ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS),
_ => ServerSideEncryption::from_static(ServerSideEncryption::AES256),
})
})
})
});
let mut effective_kms_key_id = ssekms_key_id.or_else(|| {
bucket_sse_config.as_ref().and_then(|(config, _timestamp)| {
config.rules.first().and_then(|rule| {
rule.apply_server_side_encryption_by_default
.as_ref()
.and_then(|sse| sse.kms_master_key_id.clone())
})
})
});
validate_sse_headers_for_write(
effective_sse.as_ref(),
effective_kms_key_id.as_ref(),
sse_customer_algorithm.as_ref(),
sse_customer_key.as_ref(),
sse_customer_key_md5.as_ref(),
true,
)?;
let encryption_enabled_for_put = effective_sse.is_some()
|| effective_kms_key_id.is_some()
|| sse_customer_algorithm.is_some()
|| sse_customer_key.is_some()
|| sse_customer_key_md5.is_some();
let body_plan = plan_put_object_body_with_transforms(
size,
&request_context.headers,
&key,
get_buffer_size_opt_in(size),
encryption_enabled_for_put,
);
if body_plan.ingress.kind == PutObjectIngressKind::ReducedCopyCandidate {
rustfs_io_metrics::record_put_object_attempted_fast_path(size);
debug!(
encryption_enabled = encryption_enabled_for_put,
compressed = body_plan.should_compress(),
"Zero-copy write enabled for {} byte object (bucket={}, key={})",
size,
bucket,
key
);
} else if let Some(reason) = resolve_put_transformed_fallback_reason(
body_plan.ingress.kind,
body_plan.should_compress(),
encryption_enabled_for_put,
) {
rustfs_io_metrics::record_io_fallback(rustfs_io_metrics::IoStage::PutTransform, reason);
rustfs_io_metrics::record_put_fallback(size, reason);
}
let mut metadata = metadata.unwrap_or_default();
apply_put_request_metadata(
&mut metadata,
&request_context.headers,
&key,
cache_control,
content_disposition,
content_encoding,
content_language,
content_type,
expires,
website_redirect_location,
tagging,
storage_class.clone(),
)?;
let mut opts: ObjectOptions = put_opts(&bucket, &key, version_id.clone(), &request_context.headers, metadata.clone())
.await
.map_err(ApiError::from)?;
apply_put_request_object_lock_opts(
&bucket,
object_lock_legal_hold_status,
object_lock_mode,
object_lock_retain_until_date,
&mut opts,
)
.await?;
let eager_max_bytes =
resolved_small_put_eager_max_bytes(topology_aware_small_put_eager_max_bytes(&store, opts.versioned));
let can_use_small_put_eager =
!request_uses_trailing_checksum(&request_context.headers, &request_context.trailing_headers);
let current_opts: ObjectOptions = get_opts(&bucket, &key, version_id.clone(), None, &request_context.headers)
.await
.map_err(ApiError::from)?;
match store.get_object_info(&bucket, &key, &current_opts).await {
Ok(existing_obj_info) => validate_existing_object_lock_for_write(&existing_obj_info)?,
Err(err) => {
if !is_err_object_not_found(&err) && !is_err_version_not_found(&err) {
return Err(ApiError::from(err).into());
}
}
}
let actual_size = size;
let mut hash_values = PutObjectLegacyHashValues {
md5hex: if let Some(base64_md5) = content_md5 {
let md5 = base64_simd::STANDARD
.decode_to_vec(base64_md5.as_bytes())
.map_err(|e| ApiError::from(StorageError::other(format!("Invalid content MD5: {e}"))))?;
Some(hex_simd::encode_to_string(&md5, hex_simd::AsciiCase::Lower))
} else {
None
},
sha256hex: get_content_sha256_with_query(&request_context.headers, request_context.uri_query.as_deref()),
};
let reader_stage_start = std::time::Instant::now();
let stage = if can_use_small_put_eager
&& should_use_small_put_eager_path(size, eager_max_bytes, body_plan.should_compress(), encryption_enabled_for_put)
{
small_object_eager_stage = true;
debug!(
"Plain PUT is using the eager small-object path (bucket={}, key={}, size={}, eager_max={})",
bucket, key, size, eager_max_bytes
);
build_small_put_eager_hash_stage(
body,
size,
bytes_pool.clone(),
hash_values,
&request_context.headers,
request_context.trailing_headers.clone(),
)
.await?
} else if body_plan.should_compress() {
transform_stage.mark_compression();
let algorithm = CompressionAlgorithm::default();
insert_str(&mut metadata, SUFFIX_COMPRESSION, algorithm.to_string());
insert_str(&mut metadata, SUFFIX_ACTUAL_SIZE, size.to_string());
let ingress_source = build_put_object_ingress_source(body, body_plan);
let stage = build_put_object_plain_hash_stage(
ingress_source,
std::mem::take(&mut hash_values),
PutObjectLegacyHashStagePlan {
size,
actual_size: size,
apply_s3_checksum: true,
ignore_s3_checksum_value: false,
},
&request_context.headers,
request_context.trailing_headers.clone(),
)
.map_err(ApiError::from)?;
if stage.ingress_kind == PutObjectIngressKind::ReducedCopyCandidate {
plain_reduced_copy_stage = true;
}
opts.want_checksum = stage.want_checksum;
insert_str(&mut opts.user_defined, SUFFIX_COMPRESSION, algorithm.to_string());
insert_str(&mut opts.user_defined, SUFFIX_ACTUAL_SIZE, size.to_string());
let reader: Box<dyn Reader> = Box::new(CompressReader::new(stage.reader, algorithm));
size = HashReader::SIZE_PRESERVE_LAYER;
hash_values.clear_for_transformed_body();
build_put_object_legacy_hash_stage(
reader,
hash_values,
PutObjectLegacyHashStagePlan {
size,
actual_size,
apply_s3_checksum: size >= 0,
ignore_s3_checksum_value: false,
},
&request_context.headers,
request_context.trailing_headers.clone(),
)
.map_err(ApiError::from)?
} else {
let ingress_source = build_put_object_ingress_source(body, body_plan);
let stage = build_put_object_plain_hash_stage(
ingress_source,
hash_values,
PutObjectLegacyHashStagePlan {
size,
actual_size,
apply_s3_checksum: size >= 0,
ignore_s3_checksum_value: false,
},
&request_context.headers,
request_context.trailing_headers.clone(),
)
.map_err(ApiError::from)?;
if stage.ingress_kind == PutObjectIngressKind::ReducedCopyCandidate {
plain_reduced_copy_stage = true;
debug!(
"Plain PUT is using the reduced-copy Reader + BlockReadable hash path (bucket={}, key={})",
bucket, key
);
}
stage
};
log_put_flow_phase(
&bucket,
&key,
"build_hash_stage",
reader_stage_start.elapsed(),
actual_size,
small_object_eager_stage,
plain_reduced_copy_stage,
transform_stage.compression_applied(),
false,
);
let mut reader = stage.reader;
if stage.want_checksum.is_some() {
opts.want_checksum = stage.want_checksum;
}
let encryption_request = EncryptionRequest {
bucket: &bucket,
key: &key,
server_side_encryption: effective_sse.clone(),
ssekms_key_id: effective_kms_key_id.clone(),
sse_customer_algorithm: sse_customer_algorithm.clone(),
sse_customer_key,
sse_customer_key_md5: sse_customer_key_md5.clone(),
content_size: actual_size,
part_number: None,
part_key: None,
part_nonce: None,
};
if let Some(material) = sse_encryption(encryption_request).await? {
transform_stage.mark_encryption();
effective_sse = Some(material.server_side_encryption.clone());
effective_kms_key_id = material.kms_key_id.clone();
let encrypted_reader = material.wrap_reader(reader);
reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
let encryption_metadata = material.metadata;
metadata.extend(encryption_metadata.clone());
opts.user_defined.extend(encryption_metadata);
}
let mut reader = ChunkNativePutData::new(reader);
let mt2 = metadata.clone();
opts.user_defined.extend(metadata);
let repoptions =
get_must_replicate_options(&mt2, "".to_string(), ReplicationStatusType::Empty, ReplicationType::Object, opts.clone());
let dsc = must_replicate(&bucket, &key, repoptions).await;
if dsc.replicate_any() {
insert_str(&mut opts.user_defined, SUFFIX_REPLICATION_TIMESTAMP, jiff::Zoned::now().to_string());
insert_str(
&mut opts.user_defined,
SUFFIX_REPLICATION_STATUS,
dsc.pending_status().unwrap_or_default(),
);
}
let store_put_start = std::time::Instant::now();
let obj_info = store
.put_object(&bucket, &key, &mut reader, &opts)
.await
.map_err(ApiError::from)?;
log_put_flow_phase(
&bucket,
&key,
"store_put_object",
store_put_start.elapsed(),
actual_size,
small_object_eager_stage,
plain_reduced_copy_stage,
transform_stage.compression_applied(),
transform_stage.encryption_applied(),
);
maybe_enqueue_transition_immediate(&obj_info, LcEventSrc::S3PutObject).await;
rustfs_ecstore::data_usage::increment_bucket_usage_memory(&bucket, obj_info.size as u64).await;
let raw_version = obj_info.version_id.map(|v| v.to_string());
Self::spawn_cache_invalidation(bucket.clone(), key.clone(), raw_version.clone());
let put_version = if bucket_prefix_versioning_enabled(&bucket, &key).await {
raw_version.clone()
} else {
None
};
let e_tag = obj_info.etag.clone().map(|etag| to_s3s_etag(&etag));
let repoptions =
get_must_replicate_options(&mt2, "".to_string(), ReplicationStatusType::Empty, ReplicationType::Object, opts);
let dsc = must_replicate(&bucket, &key, repoptions).await;
let expiration = resolve_put_object_expiration(&bucket, &obj_info).await;
if dsc.replicate_any() {
schedule_replication(obj_info.clone(), store.clone(), dsc, ReplicationType::Object).await;
}
let mut checksums = PutObjectChecksums {
crc32: input.checksum_crc32,
crc32c: input.checksum_crc32c,
sha1: input.checksum_sha1,
sha256: input.checksum_sha256,
crc64nvme: input.checksum_crc64nvme,
};
apply_trailing_checksums(
input.checksum_algorithm.as_ref().map(|a| a.as_str()),
&request_context.trailing_headers,
&mut checksums,
);
checksums.merge_from_map(&reader.content_crc());
if let Some(checksum_bytes) = resolved_checksum_bytes(&checksums)
&& obj_info
.checksum
.as_ref()
.is_none_or(|stored| rustfs_rio::read_checksums(stored.as_ref(), 0).0.is_empty())
{
let checksum_update_opts = ObjectOptions {
version_id: raw_version.clone(),
resolved_checksum: Some(checksum_bytes),
..Default::default()
};
let _ = store
.put_object_metadata(&bucket, &key, &checksum_update_opts)
.await
.map_err(ApiError::from)?;
}
let output = PutObjectOutput {
e_tag,
server_side_encryption: effective_sse,
sse_customer_algorithm: sse_customer_algorithm.clone(),
sse_customer_key_md5: sse_customer_key_md5.clone(),
ssekms_key_id: effective_kms_key_id,
expiration,
checksum_crc32: checksums.crc32,
checksum_crc32c: checksums.crc32c,
checksum_sha1: checksums.sha1,
checksum_sha256: checksums.sha256,
checksum_crc64nvme: checksums.crc64nvme,
version_id: put_version,
..Default::default()
};
let manager = get_capacity_manager();
manager.record_write_operation().await;
{
let duration_ms = start_time.elapsed().as_millis() as f64;
let fast_path_selected = plain_reduced_copy_stage || small_object_eager_stage;
rustfs_io_metrics::record_put_object(duration_ms, size, fast_path_selected);
let io_path = if fast_path_selected {
rustfs_io_metrics::IoPath::Fast
} else {
rustfs_io_metrics::IoPath::Legacy
};
rustfs_io_metrics::record_io_path_selected("put", io_path);
rustfs_io_metrics::record_put_path_selected(actual_size, io_path);
let effective_copy_mode = transform_stage.effective_copy_mode();
rustfs_io_metrics::record_io_copy_mode("put", effective_copy_mode, actual_size.max(0) as usize);
rustfs_io_metrics::record_put_copy_mode(actual_size, effective_copy_mode);
if let Some(transform_kind) = transform_stage.metric_kind() {
rustfs_io_metrics::record_put_transform_selected(transform_kind, io_path, actual_size.max(0) as usize);
}
}
Ok(PutObjectFlowResult {
output,
helper_object: obj_info,
helper_version_id: raw_version,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use bytes::Bytes;
use futures::{StreamExt, stream};
use rustfs_io_core::BytesPool;
use serial_test::serial;
use std::sync::Arc;
use tokio::time::{Duration, timeout};
#[test]
fn small_put_eager_path_only_targets_plain_small_objects() {
assert!(should_use_small_put_eager_path(
64 * 1024,
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES,
false,
false
));
assert!(should_use_small_put_eager_path(
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES,
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES,
false,
false
));
assert!(!should_use_small_put_eager_path(
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES + 1,
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES,
false,
false
));
assert!(!should_use_small_put_eager_path(
64 * 1024,
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES,
true,
false
));
assert!(!should_use_small_put_eager_path(
64 * 1024,
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES,
false,
true
));
assert!(!should_use_small_put_eager_path(0, DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, false, false));
}
#[test]
fn clamp_small_put_eager_max_bytes_caps_inline_budget() {
assert_eq!(
clamp_small_put_eager_max_bytes(Some(rustfs_object_io::put::PUT_REDUCED_COPY_MIN_SIZE_BYTES as usize * 2)),
DEFAULT_SMALL_PUT_EAGER_MAX_BYTES
);
assert_eq!(clamp_small_put_eager_max_bytes(Some(128 * 1024)), 128 * 1024);
assert_eq!(clamp_small_put_eager_max_bytes(None), DEFAULT_SMALL_PUT_EAGER_MAX_BYTES);
}
#[test]
#[serial]
fn resolved_small_put_eager_max_bytes_honors_disable_env() {
temp_env::with_var(ENV_RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER, Some("true"), || {
assert_eq!(resolved_small_put_eager_max_bytes(256 * 1024), 0);
});
}
#[test]
#[serial]
fn resolved_small_put_eager_max_bytes_narrows_default_budget() {
temp_env::with_var(ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES, Some("4096"), || {
assert_eq!(resolved_small_put_eager_max_bytes(256 * 1024), 4096);
});
}
#[test]
#[serial]
fn resolved_small_put_eager_max_bytes_ignores_invalid_override() {
temp_env::with_var(ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES, Some("invalid"), || {
assert_eq!(resolved_small_put_eager_max_bytes(256 * 1024), 256 * 1024);
});
}
#[tokio::test]
async fn read_small_put_body_eager_requires_exact_content_length() {
let body = stream::iter(vec![
Ok::<Bytes, std::io::Error>(Bytes::from_static(b"abc")),
Ok::<Bytes, std::io::Error>(Bytes::from_static(b"def")),
]);
let pool = Arc::new(BytesPool::new_tiered());
let data = read_small_put_body_eager(body, 6, pool)
.await
.expect("eager read should succeed");
assert_eq!(data.as_ref(), b"abcdef");
}
#[tokio::test]
async fn read_small_put_body_eager_rejects_length_mismatch() {
let body = stream::iter(vec![Ok::<Bytes, std::io::Error>(Bytes::from_static(b"abc"))]);
let pool = Arc::new(BytesPool::new_tiered());
let err = read_small_put_body_eager(body, 4, pool)
.await
.expect_err("short eager read should fail");
assert_eq!(err.code(), &S3ErrorCode::IncompleteBody);
}
#[tokio::test]
async fn read_small_put_body_eager_rejects_overlong_body() {
let body = stream::iter(vec![
Ok::<Bytes, std::io::Error>(Bytes::from_static(b"abc")),
Ok::<Bytes, std::io::Error>(Bytes::from_static(b"def")),
]);
let pool = Arc::new(BytesPool::new_tiered());
let err = read_small_put_body_eager(body, 5, pool)
.await
.expect_err("overlong eager read should fail");
assert_eq!(err.code(), &S3ErrorCode::IncompleteBody);
}
#[tokio::test]
async fn read_small_put_body_eager_returns_buffer_to_pool_after_drop() {
let body = stream::iter(vec![Ok::<Bytes, std::io::Error>(Bytes::from_static(b"abc"))]);
let pool = Arc::new(BytesPool::new_tiered());
let data = read_small_put_body_eager(body, 3, pool.clone())
.await
.expect("pooled eager read should succeed");
assert_eq!(pool.available_buffers(), 0);
drop(data);
assert_eq!(pool.available_buffers(), 1);
}
#[tokio::test]
async fn read_small_put_body_eager_returns_after_expected_bytes_without_waiting_for_eof() {
let body = stream::once(async { Ok::<Bytes, std::io::Error>(Bytes::from_static(b"abc")) }).chain(stream::pending());
let pool = Arc::new(BytesPool::new_tiered());
let data = timeout(Duration::from_millis(50), read_small_put_body_eager(body, 3, pool))
.await
.expect("eager read should not wait for stream termination")
.expect("eager read should succeed once content-length bytes are read");
assert_eq!(data.as_ref(), b"abc");
}
}
+52
View File
@@ -0,0 +1,52 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use super::*;
#[derive(Clone)]
pub(super) struct GetObjectRequestContext {
pub(super) bucket: String,
pub(super) key: String,
pub(super) cache_key: String,
pub(super) version_id_for_event: String,
pub(super) part_number: Option<usize>,
pub(super) rs: Option<HTTPRangeSpec>,
pub(super) opts: ObjectOptions,
pub(super) headers: HeaderMap,
pub(super) method: hyper::Method,
pub(super) sse_customer_key: Option<String>,
pub(super) sse_customer_key_md5: Option<String>,
}
pub(super) type PutObjectChecksums = rustfs_object_io::put::PutObjectChecksums;
#[derive(Clone)]
pub(super) struct PutObjectRequestContext {
pub(super) headers: HeaderMap,
pub(super) trailing_headers: Option<s3s::TrailingHeaders>,
pub(super) uri_query: Option<String>,
pub(super) is_post_object: bool,
pub(super) method: hyper::Method,
pub(super) uri: hyper::Uri,
pub(super) extensions: http::Extensions,
pub(super) credentials: Option<s3s::auth::Credentials>,
pub(super) region: Option<s3s::region::Region>,
pub(super) service: Option<String>,
}
pub(super) struct PutObjectFlowResult {
pub(super) output: PutObjectOutput,
pub(super) helper_object: ObjectInfo,
pub(super) helper_version_id: Option<String>,
}
File diff suppressed because it is too large Load Diff
+14
View File
@@ -259,6 +259,9 @@ impl From<std::io::Error> for ApiError {
fn from(err: std::io::Error) -> Self {
// Check if the error is a ChecksumMismatch (BadDigest)
if let Some(inner) = err.get_ref() {
if let Some(storage_error) = inner.downcast_ref::<StorageError>() {
return storage_error.clone().into();
}
if inner.downcast_ref::<rustfs_rio::ChecksumMismatch>().is_some() {
return ApiError {
code: S3ErrorCode::BadDigest,
@@ -552,4 +555,15 @@ mod tests {
// This is expected because ApiError is not a typical Error implementation
assert!(error.source().is_none());
}
#[test]
fn test_api_error_from_io_error_unwraps_invalid_range_storage_error() {
let io_error = std::io::Error::from(StorageError::InvalidRangeSpec("range invalid".to_string()));
let api_error: ApiError = io_error.into();
assert_eq!(api_error.code, S3ErrorCode::InvalidRange);
assert_eq!(api_error.message, ApiError::error_code_to_message(&S3ErrorCode::InvalidRange));
assert!(api_error.source.is_some());
}
}
+9 -6
View File
@@ -43,7 +43,7 @@ use crate::init::init_webdav_system;
use crate::capacity::capacity_integration::init_capacity_management;
use crate::server::{
SHUTDOWN_TIMEOUT, ServiceState, ServiceStateManager, ShutdownSignal, init_cert, init_event_notifier, shutdown_event_notifier,
SHUTDOWN_TIMEOUT, ServiceState, ServiceStateManager, ShutdownSignal, init_event_notifier, shutdown_event_notifier,
start_audit_system, start_http_server, stop_audit_system, wait_for_shutdown,
};
use license::{current_license, init_license, license_status};
@@ -229,15 +229,18 @@ async fn async_main() -> Result<()> {
// A crypto provider is already installed (e.g. by the host process); this is fine.
debug!("rustls crypto provider already installed, skipping aws-lc-rs default install");
}
// Initialize TLS if a certificate path is provided
// Initialize TLS outbound material (root CAs, mTLS identity) if configured.
// Server-side TLS acceptor is built separately inside start_http_server()
// using the same TlsMaterialSnapshot loading logic.
if let Some(tls_path) = &config.tls_path {
match init_cert(tls_path).await {
Ok(_) => {
info!(target: "rustfs::main", "TLS initialized successfully with certs from {}", tls_path);
match crate::server::tls_material::TlsMaterialSnapshot::load(tls_path).await {
Ok(snapshot) => {
snapshot.apply_outbound().await;
info!(target: "rustfs::main", "TLS outbound material initialized from {}", tls_path);
}
Err(e) => {
error!("Failed to initialize TLS from {}: {}", tls_path, e);
return Err(Error::other(e));
return Err(Error::other(e.to_string()));
}
}
}
+1
View File
@@ -81,6 +81,7 @@ impl ProtocolStorageClient {
object: params.object,
version_id: None,
region: None,
request_context: Some(crate::storage::request_context::RequestContext::fallback()),
});
let req = S3Request {
+18 -8
View File
@@ -14,13 +14,27 @@
use crate::app::context::resolve_server_config;
use rustfs_audit::{AuditError, AuditResult, audit_system, init_audit_system, system::AuditSystemState};
use rustfs_config::DEFAULT_DELIMITER;
use tracing::{info, warn};
fn server_config_from_context() -> Option<rustfs_ecstore::config::Config> {
resolve_server_config()
}
fn has_any_audit_targets(config: &rustfs_ecstore::config::Config) -> bool {
for subsystem in [
rustfs_config::audit::AUDIT_MQTT_SUB_SYS,
rustfs_config::audit::AUDIT_WEBHOOK_SUB_SYS,
] {
let Some(targets) = config.0.get(subsystem) else {
continue;
};
if targets.keys().any(|key| key != rustfs_config::DEFAULT_DELIMITER) {
return true;
}
}
false
}
/// Start the audit system.
/// This function checks if the audit subsystem is configured in the global server configuration.
/// If configured, it initializes and starts the audit system.
@@ -55,10 +69,8 @@ pub(crate) async fn start_audit_system() -> AuditResult<()> {
"The global server configuration is loaded"
);
// 2. Check if the notify subsystem exists in the configuration, and skip initialization if it doesn't
let mqtt_config = server_config.get_value(rustfs_config::audit::AUDIT_MQTT_SUB_SYS, DEFAULT_DELIMITER);
let webhook_config = server_config.get_value(rustfs_config::audit::AUDIT_WEBHOOK_SUB_SYS, DEFAULT_DELIMITER);
if mqtt_config.is_none() && webhook_config.is_none() {
let has_targets = has_any_audit_targets(&server_config);
if !has_targets {
info!(
target: "rustfs::main::start_audit_system",
"Audit subsystem (MQTT/Webhook) is not configured, and audit system initialization is skipped."
@@ -68,9 +80,7 @@ pub(crate) async fn start_audit_system() -> AuditResult<()> {
info!(
target: "rustfs::main::start_audit_system",
"Audit subsystem configuration detected (MQTT: {}, Webhook: {}) and started initializing the audit system.",
mqtt_config.is_some(),
webhook_config.is_some()
"Audit subsystem configuration detected and started initializing the audit system."
);
// 3. Initialize and start the audit system
let system = init_audit_system();
-296
View File
@@ -1,296 +0,0 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use rustfs_common::{MtlsIdentityPem, set_global_mtls_identity, set_global_root_cert};
use rustfs_config::{RUSTFS_CA_CERT, RUSTFS_PUBLIC_CERT, RUSTFS_TLS_CERT};
use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject};
use std::path::{Path, PathBuf};
use tracing::{debug, info};
#[derive(Debug)]
pub enum RustFSError {
Cert(String),
}
impl std::fmt::Display for RustFSError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RustFSError::Cert(msg) => write!(f, "Certificate error: {msg}"),
}
}
}
impl std::error::Error for RustFSError {}
/// Parse PEM-encoded certificates into DER format.
/// Returns a vector of DER-encoded certificates.
///
/// # Arguments
/// * `pem` - A byte slice containing the PEM-encoded certificates.
///
/// # Returns
/// A vector of `CertificateDer` containing the DER-encoded certificates.
///
/// # Errors
/// Returns `RustFSError` if parsing fails.
fn parse_pem_certs(pem: &[u8]) -> Result<Vec<CertificateDer<'static>>, RustFSError> {
let mut out = Vec::new();
let mut reader = std::io::Cursor::new(pem);
for item in CertificateDer::pem_reader_iter(&mut reader) {
let c = item.map_err(|e| RustFSError::Cert(format!("parse cert pem: {e}")))?;
out.push(c);
}
Ok(out)
}
/// Parse a PEM-encoded private key into DER format.
/// Supports PKCS#8 and RSA private keys.
///
/// # Arguments
/// * `pem` - A byte slice containing the PEM-encoded private key.
///
/// # Returns
/// A `PrivateKeyDer` containing the DER-encoded private key.
///
/// # Errors
/// Returns `RustFSError` if parsing fails or no key is found.
fn parse_pem_private_key(pem: &[u8]) -> Result<PrivateKeyDer<'static>, RustFSError> {
let mut reader = std::io::Cursor::new(pem);
PrivateKeyDer::from_pem_reader(&mut reader).map_err(|e| RustFSError::Cert(format!("parse private key pem: {e}")))
}
/// Helper function to read a file and return its contents.
/// Returns the file contents as a vector of bytes.
/// # Errors
/// Returns `RustFSError` if reading fails.
async fn read_file(path: &PathBuf, desc: &str) -> Result<Vec<u8>, RustFSError> {
tokio::fs::read(path)
.await
.map_err(|e| RustFSError::Cert(format!("read {desc} {path:?}: {e}")))
}
/// Initialize TLS material for both server and outbound client connections.
///
/// Loads roots from:
/// - `${RUSTFS_TLS_PATH}/ca.crt` (or `tls/ca.crt`)
/// - `${RUSTFS_TLS_PATH}/public.crt` (optional additional root bundle)
/// - system roots if `RUSTFS_TRUST_SYSTEM_CA=true` (default: false)
/// - if `RUSTFS_TRUST_LEAF_CERT_AS_CA=true`, also loads leaf cert(s) from
/// `${RUSTFS_TLS_PATH}/rustfs_cert.pem` into the root store.
///
/// Loads mTLS client identity (optional) from:
/// - `${RUSTFS_TLS_PATH}/client_cert.pem`
/// - `${RUSTFS_TLS_PATH}/client_key.pem`
///
/// Environment overrides:
/// - RUSTFS_TLS_PATH
/// - RUSTFS_MTLS_CLIENT_CERT
/// - RUSTFS_MTLS_CLIENT_KEY
pub(crate) async fn init_cert(tls_path: &str) -> Result<(), RustFSError> {
if tls_path.is_empty() {
info!("No TLS path configured; skipping certificate initialization");
return Ok(());
}
let tls_dir = PathBuf::from(tls_path);
// Load root certificates
load_root_certs(&tls_dir).await?;
// Load optional mTLS identity
load_mtls_identity(&tls_dir).await?;
Ok(())
}
/// Load root certificates from various sources.
async fn load_root_certs(tls_dir: &Path) -> Result<(), RustFSError> {
let mut cert_data = Vec::new();
let trust_leaf_as_ca =
rustfs_utils::get_env_bool(rustfs_config::ENV_TRUST_LEAF_CERT_AS_CA, rustfs_config::DEFAULT_TRUST_LEAF_CERT_AS_CA);
if trust_leaf_as_ca {
walk_dir(tls_dir.to_path_buf(), RUSTFS_TLS_CERT, &mut cert_data).await;
info!("Loaded leaf certificate(s) as root CA as per RUSTFS_TRUST_LEAF_CERT_AS_CA");
}
// Try public.crt and ca.crt
let public_cert_path = tls_dir.join(RUSTFS_PUBLIC_CERT);
load_cert_file(public_cert_path.to_str().unwrap_or_default(), &mut cert_data, "CA certificate").await;
let ca_cert_path = tls_dir.join(RUSTFS_CA_CERT);
load_cert_file(ca_cert_path.to_str().unwrap_or_default(), &mut cert_data, "CA certificate").await;
// Load system root certificates if enabled
let trust_system_ca = rustfs_utils::get_env_bool(rustfs_config::ENV_TRUST_SYSTEM_CA, rustfs_config::DEFAULT_TRUST_SYSTEM_CA);
if trust_system_ca {
let system_ca_paths = [
"/etc/ssl/certs/ca-certificates.crt", // Debian/Ubuntu/Alpine
"/etc/pki/tls/certs/ca-bundle.crt", // Fedora/RHEL/CentOS
"/etc/ssl/ca-bundle.pem", // OpenSUSE
"/etc/pki/tls/cacert.pem", // OpenELEC
"/etc/ssl/cert.pem", // macOS/FreeBSD
"/usr/local/etc/openssl/cert.pem", // macOS/Homebrew OpenSSL
"/usr/local/share/certs/ca-root-nss.crt", // FreeBSD
"/etc/pki/ca-trust/extracted/pem/tls-ca-bundle.pem", // RHEL
"/usr/share/pki/ca-trust-legacy/ca-bundle.legacy.crt", // RHEL legacy
];
let mut system_cert_loaded = false;
for path in system_ca_paths {
if load_cert_file(path, &mut cert_data, "system root certificates").await {
system_cert_loaded = true;
info!("Loaded system root certificates from {}", path);
break;
}
}
if !system_cert_loaded {
debug!("Could not find system root certificates in common locations.");
}
} else {
info!("Loading system root certificates disabled via RUSTFS_TRUST_SYSTEM_CA");
}
if !cert_data.is_empty() {
set_global_root_cert(cert_data).await;
info!("Configured custom root certificates for inter-node communication");
}
Ok(())
}
/// Load optional mTLS identity.
async fn load_mtls_identity(tls_dir: &Path) -> Result<(), RustFSError> {
let client_cert_path = match rustfs_utils::get_env_opt_str(rustfs_config::ENV_MTLS_CLIENT_CERT) {
Some(p) => PathBuf::from(p),
None => tls_dir.join(rustfs_config::RUSTFS_CLIENT_CERT_FILENAME),
};
let client_key_path = match rustfs_utils::get_env_opt_str(rustfs_config::ENV_MTLS_CLIENT_KEY) {
Some(p) => PathBuf::from(p),
None => tls_dir.join(rustfs_config::RUSTFS_CLIENT_KEY_FILENAME),
};
if client_cert_path.exists() && client_key_path.exists() {
let cert_bytes = read_file(&client_cert_path, "client cert").await?;
let key_bytes = read_file(&client_key_path, "client key").await?;
// Validate parse-ability early; store as PEM bytes for tonic.
parse_pem_certs(&cert_bytes)?;
parse_pem_private_key(&key_bytes)?;
let identity_pem = MtlsIdentityPem {
cert_pem: cert_bytes,
key_pem: key_bytes,
};
set_global_mtls_identity(Some(identity_pem)).await;
info!("Loaded mTLS client identity cert={:?} key={:?}", client_cert_path, client_key_path);
} else {
set_global_mtls_identity(None).await;
info!(
"mTLS client identity not configured (missing {:?} and/or {:?}); proceeding with server-only TLS",
client_cert_path, client_key_path
);
}
Ok(())
}
/// Helper function to load a certificate file and append to cert_data.
/// Returns true if the file was successfully loaded.
async fn load_cert_file(path: &str, cert_data: &mut Vec<u8>, desc: &str) -> bool {
if tokio::fs::metadata(path).await.is_ok() {
if let Ok(data) = tokio::fs::read(path).await {
cert_data.extend(data);
cert_data.push(b'\n');
info!("Loaded {} from {}", desc, path);
true
} else {
debug!("Failed to read {} from {}", desc, path);
false
}
} else {
debug!("{} file not found at {}", desc, path);
false
}
}
/// Load the certificate file if its name matches `cert_name`.
/// If it matches, the certificate data is appended to `cert_data`.
///
/// # Parameters
/// - `entry`: The directory entry to check.
/// - `cert_name`: The name of the certificate file to match.
/// - `cert_data`: A mutable vector to append loaded certificate data.
async fn load_if_matches(entry: &tokio::fs::DirEntry, cert_name: &str, cert_data: &mut Vec<u8>) {
let fname = entry.file_name().to_string_lossy().to_string();
if fname == cert_name {
let p = entry.path();
load_cert_file(&p.to_string_lossy(), cert_data, "certificate").await;
}
}
/// Search the directory at `path` and one level of subdirectories to find and load
/// certificates matching `cert_name`. Loaded certificate data is appended to
/// `cert_data`.
/// # Parameters
/// - `path`: The starting directory path to search for certificates.
/// - `cert_name`: The name of the certificate file to look for.
/// - `cert_data`: A mutable vector to append loaded certificate data.
async fn walk_dir(path: PathBuf, cert_name: &str, cert_data: &mut Vec<u8>) {
if let Ok(mut rd) = tokio::fs::read_dir(&path).await {
while let Ok(Some(entry)) = rd.next_entry().await {
if let Ok(ft) = entry.file_type().await {
if ft.is_file() {
load_if_matches(&entry, cert_name, cert_data).await;
} else if ft.is_dir() {
// Only check direct subdirectories, no deeper recursion
if let Ok(mut sub_rd) = tokio::fs::read_dir(&entry.path()).await {
while let Ok(Some(sub_entry)) = sub_rd.next_entry().await {
if let Ok(sub_ft) = sub_entry.file_type().await
&& sub_ft.is_file()
{
load_if_matches(&sub_entry, cert_name, cert_data).await;
}
// Ignore subdirectories and symlinks in subdirs to limit to one level
}
}
} else if ft.is_symlink() {
// Follow symlink and treat target as file or directory, but limit to one level
if let Ok(meta) = tokio::fs::metadata(&entry.path()).await {
if meta.is_file() {
load_if_matches(&entry, cert_name, cert_data).await;
} else if meta.is_dir() {
// Treat as directory but only check its direct contents
if let Ok(mut sub_rd) = tokio::fs::read_dir(&entry.path()).await {
while let Ok(Some(sub_entry)) = sub_rd.next_entry().await {
if let Ok(sub_ft) = sub_entry.file_type().await
&& sub_ft.is_file()
{
load_if_matches(&sub_entry, cert_name, cert_data).await;
}
// Ignore deeper levels
}
}
}
}
}
}
}
} else {
debug!("Certificate directory not found: {}", path.display());
}
}
+196
View File
@@ -50,6 +50,11 @@ use std::str::FromStr;
use tower_http::compression::predicate::Predicate;
use tracing::debug;
/// Response extension key for storing the request path category.
/// Set by `PathCategoryInjectionLayer` before the compression predicate evaluates.
#[derive(Debug, Clone, Copy)]
pub(crate) struct RequestPathCategory(pub(crate) PathCategory);
/// Configuration for HTTP response compression.
///
/// This structure holds the whitelist-based compression settings:
@@ -319,6 +324,156 @@ impl Predicate for CompressionPredicate {
}
}
// ── Path-Aware Compression ──
/// Classifies request paths to determine if compression should apply.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum PathCategory {
/// S3 data plane (bucket/key operations) — compression applies via whitelist
S3DataPlane,
/// Admin API paths — skip compression (small JSON responses)
AdminApi,
/// Console paths — skip compression (static assets, already optimized)
Console,
/// Internode RPC paths — skip compression (binary protocol data)
InternodeRpc,
/// Health/probe paths — skip compression (tiny responses)
Probe,
}
impl PathCategory {
/// Classify a request URI path into a category.
pub(crate) fn classify(path: &str) -> Self {
if path.starts_with("/rustfs/rpc/") || path.starts_with("/rustfs/peer/") {
PathCategory::InternodeRpc
} else if path.starts_with("/rustfs/admin/") || path.starts_with("/minio/admin/") {
PathCategory::AdminApi
} else if path.starts_with("/rustfs/console") {
PathCategory::Console
} else if path.starts_with("/minio/health/") {
PathCategory::Probe
} else {
PathCategory::S3DataPlane
}
}
/// Returns true if compression should be considered for this path category.
/// Only S3 data plane paths go through the full compression predicate.
#[inline]
pub(crate) fn should_evaluate_compression(self) -> bool {
matches!(self, PathCategory::S3DataPlane)
}
}
/// A compression predicate that first checks the request path category
/// before evaluating the full compression rules.
///
/// This avoids running MIME type / extension matching for admin, RPC, console,
/// and health probe paths where compression is never beneficial.
#[derive(Clone, Debug)]
pub(crate) struct PathAwareCompressionPredicate {
inner: CompressionPredicate,
}
impl PathAwareCompressionPredicate {
pub(crate) fn new(config: CompressionConfig) -> Self {
Self {
inner: CompressionPredicate::new(config),
}
}
}
impl Predicate for PathAwareCompressionPredicate {
fn should_compress<B>(&self, response: &Response<B>) -> bool
where
B: http_body::Body,
{
// Fast path: skip full predicate evaluation for non-S3 paths
if let Some(RequestPathCategory(category)) = response.extensions().get::<RequestPathCategory>()
&& !category.should_evaluate_compression()
{
return false;
}
self.inner.should_compress(response)
}
}
use http::Request;
use http_body::Body;
use std::pin::Pin;
use std::task::{Context, Poll};
use tower::{Layer, Service};
/// Tower layer that injects `RequestPathCategory` into each response's extensions
/// based on the incoming request URI path. Must be placed before `CompressionLayer`.
#[derive(Clone, Copy, Debug)]
pub(crate) struct PathCategoryInjectionLayer;
impl<S> Layer<S> for PathCategoryInjectionLayer {
type Service = PathCategoryInjectionService<S>;
fn layer(&self, inner: S) -> Self::Service {
PathCategoryInjectionService { inner }
}
}
/// Service wrapper that adds `RequestPathCategory` to response extensions.
#[derive(Clone)]
pub(crate) struct PathCategoryInjectionService<S> {
inner: S,
}
pin_project_lite::pin_project! {
/// Future for `PathCategoryInjectionService` that injects path category into response.
#[project = InjectCategoryFutProj]
pub(crate) struct InjectCategoryFut<F> {
#[pin]
inner: F,
category: PathCategory,
}
}
impl<F, ResBody, E> std::future::Future for InjectCategoryFut<F>
where
F: std::future::Future<Output = Result<Response<ResBody>, E>>,
{
type Output = Result<Response<ResBody>, E>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.project();
match this.inner.poll(cx) {
Poll::Ready(Ok(mut resp)) => {
resp.extensions_mut().insert(RequestPathCategory(*this.category));
Poll::Ready(Ok(resp))
}
Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
Poll::Pending => Poll::Pending,
}
}
}
impl<S, ReqBody, ResBody> Service<Request<ReqBody>> for PathCategoryInjectionService<S>
where
S: Service<Request<ReqBody>, Response = Response<ResBody>>,
ResBody: Body,
{
type Response = Response<ResBody>;
type Error = S::Error;
type Future = InjectCategoryFut<S::Future>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
let category = PathCategory::classify(req.uri().path());
InjectCategoryFut {
inner: self.inner.call(req),
category,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
@@ -471,4 +626,45 @@ mod tests {
assert_eq!(predicate.config.mime_patterns.len(), 2);
assert_eq!(predicate.config.min_size, 1000);
}
#[test]
fn test_path_category_classify_s3() {
assert_eq!(PathCategory::classify("/"), PathCategory::S3DataPlane);
assert_eq!(PathCategory::classify("/mybucket"), PathCategory::S3DataPlane);
assert_eq!(PathCategory::classify("/mybucket/mykey"), PathCategory::S3DataPlane);
assert_eq!(PathCategory::classify("/bucket?list-type=2"), PathCategory::S3DataPlane);
}
#[test]
fn test_path_category_classify_admin() {
assert_eq!(PathCategory::classify("/rustfs/admin/v3/service"), PathCategory::AdminApi);
assert_eq!(PathCategory::classify("/minio/admin/v3/info"), PathCategory::AdminApi);
}
#[test]
fn test_path_category_classify_console() {
assert_eq!(PathCategory::classify("/rustfs/console/index.html"), PathCategory::Console);
assert_eq!(PathCategory::classify("/rustfs/console"), PathCategory::Console);
}
#[test]
fn test_path_category_classify_rpc() {
assert_eq!(PathCategory::classify("/rustfs/rpc/read_file_stream"), PathCategory::InternodeRpc);
assert_eq!(PathCategory::classify("/rustfs/peer/health"), PathCategory::InternodeRpc);
}
#[test]
fn test_path_category_classify_probe() {
assert_eq!(PathCategory::classify("/minio/health/live"), PathCategory::Probe);
assert_eq!(PathCategory::classify("/minio/health/ready"), PathCategory::Probe);
}
#[test]
fn test_path_category_should_evaluate() {
assert!(PathCategory::S3DataPlane.should_evaluate_compression());
assert!(!PathCategory::AdminApi.should_evaluate_compression());
assert!(!PathCategory::Console.should_evaluate_compression());
assert!(!PathCategory::InternodeRpc.should_evaluate_compression());
assert!(!PathCategory::Probe.should_evaluate_compression());
}
}
+247 -169
View File
@@ -19,9 +19,13 @@ use crate::auth_keystone;
use crate::config;
use crate::server::{
ReadinessGateLayer, RemoteAddr, ServiceState, ServiceStateManager,
compress::{CompressionConfig, CompressionPredicate},
compress::{CompressionConfig, PathAwareCompressionPredicate, PathCategoryInjectionLayer},
hybrid::hybrid,
layer::{AdminChunkedContentLengthCompatLayer, ConditionalCorsLayer, ObjectAttributesEtagFixLayer, RedirectLayer},
layer::{
AdminChunkedContentLengthCompatLayer, ConditionalCorsLayer, ObjectAttributesEtagFixLayer, RedirectLayer,
RequestContextLayer,
},
tls_material::{TlsAcceptorHolder, TlsHandshakeFailureKind, TlsMaterialSnapshot, spawn_reload_loop},
};
use crate::storage;
use crate::storage::rpc::InternodeRpcService;
@@ -38,7 +42,6 @@ use metrics::{counter, histogram};
use opentelemetry::global;
use opentelemetry::trace::TraceContextExt;
use rustfs_common::GlobalReadiness;
use rustfs_config::{RUSTFS_TLS_CERT, RUSTFS_TLS_KEY};
use rustfs_ecstore::rpc::{TONIC_RPC_PREFIX, verify_rpc_signature};
use rustfs_keystone::KeystoneAuthLayer;
#[cfg(feature = "swift")]
@@ -46,7 +49,6 @@ use rustfs_protocols::SwiftService;
use rustfs_protos::proto_gen::node_service::node_service_server::NodeServiceServer;
use rustfs_trusted_proxies::ClientInfo;
use rustfs_utils::net::parse_and_resolve_address;
use rustls::ServerConfig;
use s3s::{host::MultiDomain, service::S3Service, service::S3ServiceBuilder};
use socket2::{SockRef, TcpKeepalive};
use std::io::{Error, Result};
@@ -54,7 +56,6 @@ use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use tokio::net::{TcpListener, TcpStream};
use tokio_rustls::TlsAcceptor;
use tonic::{Request, Status};
use tower::ServiceBuilder;
use tower_http::add_extension::AddExtensionLayer;
@@ -156,9 +157,24 @@ pub async fn start_http_server(
TcpListener::from_std(socket.into())?
};
let tls_acceptor = setup_tls_acceptor(config.tls_path.as_deref().unwrap_or_default()).await?;
let tls_path = config.tls_path.as_deref().unwrap_or_default();
// Load TLS materials and build server acceptor.
// Note: outbound material (root CAs, mTLS identity) is already applied in main.rs.
let tls_snapshot = TlsMaterialSnapshot::load(tls_path)
.await
.map_err(|e| std::io::Error::other(e.to_string()))?;
let tls_acceptor = tls_snapshot
.build_tls_acceptor(tls_path)
.await
.map_err(|e| std::io::Error::other(e.to_string()))?;
let tls_enabled = tls_acceptor.is_some();
let protocol = if tls_enabled { "https" } else { "http" };
// Spawn background TLS certificate hot-reload loop (if enabled).
if let Some(holder) = &tls_acceptor {
spawn_reload_loop(tls_path.to_string(), holder.clone());
}
// Obtain the listener address
let local_addr: SocketAddr = listener.local_addr()?;
let local_ip = match rustfs_utils::get_local_ip() {
@@ -273,11 +289,53 @@ pub async fn start_http_server(
(sigterm_inner, sigint_inner)
};
// RustFS Transport Layer Configuration Constants - Optimized for S3 Workloads
const H2_INITIAL_STREAM_WINDOW_SIZE: u32 = 1024 * 1024 * 4; // 4MB: Optimize large file throughput
const H2_INITIAL_CONN_WINDOW_SIZE: u32 = 1024 * 1024 * 8; // 8MB: Link-level flow control
const H2_MAX_FRAME_SIZE: u32 = 512 * 1024; // 512KB: Reduce framing overhead for large objects
const H2_MAX_HEADER_LIST_SIZE: u32 = 64 * 1024; // 64KB: Conservative header limit to mitigate DoS risk
// ── HTTP Transport Tuning (configurable via env vars) ──
// Read all transport parameters from environment, falling back to defaults.
// H2 frame size is clamped to RFC 7540 range: 2^14 (16KB) to 2^24 (16MB).
let h2_stream_window = rustfs_utils::get_env_u32(
rustfs_config::ENV_H2_INITIAL_STREAM_WINDOW_SIZE,
rustfs_config::DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE,
);
let h2_conn_window = rustfs_utils::get_env_u32(
rustfs_config::ENV_H2_INITIAL_CONN_WINDOW_SIZE,
rustfs_config::DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE,
);
let h2_max_frame_size =
rustfs_utils::get_env_u32(rustfs_config::ENV_H2_MAX_FRAME_SIZE, rustfs_config::DEFAULT_H2_MAX_FRAME_SIZE)
.clamp(16_384, 16_777_216); // RFC 7540
let h2_max_header_list_size =
rustfs_utils::get_env_u32(rustfs_config::ENV_H2_MAX_HEADER_LIST_SIZE, rustfs_config::DEFAULT_H2_MAX_HEADER_LIST_SIZE);
let h2_max_concurrent_streams = rustfs_utils::get_env_u32(
rustfs_config::ENV_H2_MAX_CONCURRENT_STREAMS,
rustfs_config::DEFAULT_H2_MAX_CONCURRENT_STREAMS,
)
.max(1);
let h2_keep_alive_interval =
rustfs_utils::get_env_u64(rustfs_config::ENV_H2_KEEP_ALIVE_INTERVAL, rustfs_config::DEFAULT_H2_KEEP_ALIVE_INTERVAL);
let h2_keep_alive_timeout =
rustfs_utils::get_env_u64(rustfs_config::ENV_H2_KEEP_ALIVE_TIMEOUT, rustfs_config::DEFAULT_H2_KEEP_ALIVE_TIMEOUT);
let http1_header_read_timeout = rustfs_utils::get_env_u64(
rustfs_config::ENV_HTTP1_HEADER_READ_TIMEOUT,
rustfs_config::DEFAULT_HTTP1_HEADER_READ_TIMEOUT,
);
let http1_max_buf_size =
rustfs_utils::get_env_usize(rustfs_config::ENV_HTTP1_MAX_BUF_SIZE, rustfs_config::DEFAULT_HTTP1_MAX_BUF_SIZE);
info!(
"HTTP transport parameters: h2_stream_window={}, h2_conn_window={}, h2_max_frame={}, \
h2_max_header_list={}, h2_max_concurrent_streams={}, h2_keepalive_interval={}s, \
h2_keepalive_timeout={}s, http1_header_timeout={}s, http1_max_buf={}",
h2_stream_window,
h2_conn_window,
h2_max_frame_size,
h2_max_header_list_size,
h2_max_concurrent_streams,
h2_keep_alive_interval,
h2_keep_alive_timeout,
http1_header_read_timeout,
http1_max_buf_size,
);
let mut conn_builder = ConnBuilder::new(TokioExecutor::new());
@@ -286,8 +344,8 @@ pub async fn start_http_server(
.http1()
.timer(TokioTimer::new())
.keep_alive(true)
.header_read_timeout(Duration::from_secs(5))
.max_buf_size(64 * 1024)
.header_read_timeout(Duration::from_secs(http1_header_read_timeout))
.max_buf_size(http1_max_buf_size)
.writev(true);
// Optimize for HTTP/2 (AI/Data Lake high concurrency synchronization)
@@ -295,13 +353,13 @@ pub async fn start_http_server(
.http2()
.timer(TokioTimer::new())
.adaptive_window(true)
.initial_stream_window_size(H2_INITIAL_STREAM_WINDOW_SIZE)
.initial_connection_window_size(H2_INITIAL_CONN_WINDOW_SIZE)
.max_frame_size(H2_MAX_FRAME_SIZE)
.max_concurrent_streams(Some(2048))
.max_header_list_size(H2_MAX_HEADER_LIST_SIZE)
.keep_alive_interval(Some(Duration::from_secs(20)))
.keep_alive_timeout(Duration::from_secs(10));
.initial_stream_window_size(h2_stream_window)
.initial_connection_window_size(h2_conn_window)
.max_frame_size(h2_max_frame_size)
.max_concurrent_streams(Some(h2_max_concurrent_streams))
.max_header_list_size(h2_max_header_list_size)
.keep_alive_interval(Some(Duration::from_secs(h2_keep_alive_interval)))
.keep_alive_timeout(Duration::from_secs(h2_keep_alive_timeout));
let http_server = Arc::new(conn_builder);
let mut ctrl_c = std::pin::pin!(tokio::signal::ctrl_c());
@@ -310,10 +368,7 @@ pub async fn start_http_server(
// service ready
worker_state_manager.update(ServiceState::Ready);
let tls_acceptor = tls_acceptor.map(Arc::new);
// Initialize keepalive configuration once to avoid recreation in the loop
let keepalive_conf = get_default_tcp_keepalive();
// tls_acceptor is already Option<Arc<TlsAcceptorHolder>>, clone for the loop
loop {
debug!("Waiting for new connection...");
@@ -371,31 +426,32 @@ pub async fn start_http_server(
}
}
};
#[allow(unused)]
let socket_ref = SockRef::from(&socket);
// Enable TCP Keepalive to detect dead clients (e.g. power loss)
if let Err(err) = socket_ref.set_tcp_keepalive(&keepalive_conf) {
warn!(?err, "Failed to set TCP_KEEPALIVE");
}
// ── POST-ACCEPT SOCKET SYSCALLS ──
// The listening socket already sets TCP_NODELAY, TCP_KEEPALIVE,
// SO_RCVBUF, and SO_SNDBUF. On Linux/BSD, these are inherited by
// accepted sockets, so we skip redundant re-application here.
//
// Only TCP_QUICKACK (Linux) is kept — it is inherently per-connection
// and NOT inherited from the listening socket.
//
// T03 optimized: syscall count reduced from 5 → 1 (Linux) / 0 (other)
// Disable Nagle algorithm: Critical for 4KB Payload, achieving ultra-low latency
if let Err(err) = socket_ref.set_tcp_nodelay(true) {
warn!(?err, "Failed to set TCP_NODELAY");
}
// Enable TCP QuickAck to reduce latency for small requests
// Enable TCP QuickAck to reduce latency for small requests (Linux only)
#[cfg(target_os = "linux")]
if let Err(err) = socket_ref.set_tcp_quickack(true) {
debug!(?err, "Failed to set TCP_QUICKACK");
}
// Increase receive/send buffer to support BDP at GB-level throughput
if let Err(err) = socket_ref.set_recv_buffer_size(4 * rustfs_config::MI_B) {
warn!(?err, "Failed to set set_recv_buffer_size");
}
if let Err(err) = socket_ref.set_send_buffer_size(4 * rustfs_config::MI_B) {
warn!(?err, "Failed to set set_send_buffer_size");
// Debug-only: verify listening socket options were inherited
#[cfg(debug_assertions)]
{
debug!(
nodelay = socket_ref.tcp_nodelay().unwrap_or(false),
"TCP_NODELAY inherited from listening socket"
);
}
let connection_ctx = ConnectionContext {
@@ -404,6 +460,8 @@ pub async fn start_http_server(
compression_config: compression_config.clone(),
is_console,
readiness: readiness.clone(),
keystone_auth: auth_keystone::get_keystone_auth(),
trusted_proxy_layer: rustfs_trusted_proxies::is_enabled().then(|| rustfs_trusted_proxies::layer().clone()),
};
process_connection(socket, tls_acceptor.clone(), connection_ctx, graceful.clone());
@@ -433,88 +491,6 @@ pub async fn start_http_server(
Ok(shutdown_tx)
}
/// Sets up the TLS acceptor if certificates are available.
#[instrument(skip(tls_path))]
async fn setup_tls_acceptor(tls_path: &str) -> Result<Option<TlsAcceptor>> {
if tls_path.is_empty() || tokio::fs::metadata(tls_path).await.is_err() {
debug!("TLS path is not provided or does not exist, starting with HTTP");
return Ok(None);
}
debug!("Found TLS directory, checking for certificates");
let mtls_verifier = rustfs_utils::build_webpki_client_verifier(tls_path)?;
// 1. Attempt to load all certificates in the directory (multi-certificate support, for SNI)
if let Ok(cert_key_pairs) = rustfs_utils::load_all_certs_from_directory(tls_path)
&& !cert_key_pairs.is_empty()
{
debug!("Found {} certificates, creating SNI-aware multi-cert resolver", cert_key_pairs.len());
// Create an SNI-enabled certificate resolver
let resolver = rustfs_utils::create_multi_cert_resolver(cert_key_pairs)?;
// Configure the server to enable SNI support
let mut server_config = if let Some(verifier) = mtls_verifier.clone() {
ServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_cert_resolver(Arc::new(resolver))
} else {
ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(Arc::new(resolver))
};
// Configure ALPN protocol priority
server_config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec(), b"http/1.0".to_vec()];
// Enable session resumption to reduce handshake overhead for returning clients
server_config.session_storage = rustls::server::ServerSessionMemoryCache::new(10000);
// Log SNI requests
if rustfs_utils::tls_key_log() {
server_config.key_log = Arc::new(rustls::KeyLogFile::new());
}
return Ok(Some(TlsAcceptor::from(Arc::new(server_config))));
}
// 2. Revert to the traditional single-certificate mode
let key_path = format!("{tls_path}/{RUSTFS_TLS_KEY}");
let cert_path = format!("{tls_path}/{RUSTFS_TLS_CERT}");
if tokio::try_join!(tokio::fs::metadata(&key_path), tokio::fs::metadata(&cert_path)).is_ok() {
debug!("Found legacy single TLS certificate, starting with HTTPS");
let certs = rustfs_utils::load_certs(&cert_path).map_err(|e| rustfs_utils::certs_error(e.to_string()))?;
let key = rustfs_utils::load_private_key(&key_path).map_err(|e| rustfs_utils::certs_error(e.to_string()))?;
let mut server_config = if let Some(verifier) = mtls_verifier {
ServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_single_cert(certs, key)
.map_err(|e| rustfs_utils::certs_error(e.to_string()))?
} else {
ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.map_err(|e| rustfs_utils::certs_error(e.to_string()))?
};
// Configure ALPN protocol priority
server_config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec(), b"http/1.0".to_vec()];
// Enable session resumption to reduce handshake overhead for returning clients
server_config.session_storage = rustls::server::ServerSessionMemoryCache::new(10000);
// Log SNI requests
if rustfs_utils::tls_key_log() {
server_config.key_log = Arc::new(rustls::KeyLogFile::new());
}
return Ok(Some(TlsAcceptor::from(Arc::new(server_config))));
}
debug!("No valid TLS certificates found in the directory, starting with HTTP");
Ok(None)
}
#[derive(Clone)]
struct ConnectionContext {
http_server: Arc<ConnBuilder<TokioExecutor>>,
@@ -522,6 +498,10 @@ struct ConnectionContext {
compression_config: CompressionConfig,
is_console: bool,
readiness: Arc<GlobalReadiness>,
/// Pre-computed Keystone auth provider (avoids per-connection OnceLock read).
keystone_auth: Option<std::sync::Arc<rustfs_keystone::KeystoneAuthProvider>>,
/// Pre-computed trusted proxy layer (avoids per-connection is_enabled() check).
trusted_proxy_layer: Option<rustfs_trusted_proxies::TrustedProxyLayer>,
}
/// Adapter that implements the OpenTelemetry [`Extractor`] trait for Hyper's
@@ -569,7 +549,7 @@ impl<'a> opentelemetry::propagation::Extractor for HeaderMapCarrier<'a> {
))]
fn process_connection(
socket: TcpStream,
tls_acceptor: Option<Arc<TlsAcceptor>>,
tls_acceptor: Option<Arc<TlsAcceptorHolder>>,
context: ConnectionContext,
graceful: Arc<GracefulShutdown>,
) {
@@ -580,15 +560,16 @@ fn process_connection(
compression_config,
is_console,
readiness,
keystone_auth,
trusted_proxy_layer,
} = context;
// Build services inside each connected task to avoid passing complex service types across tasks,
// It also ensures that each connection has an independent service instance.
// Build the hybrid service per-connection.
// Note: NodeService is not Clone (holds LocalPeerS3Client), and the SwiftService
// type is feature-gated, so we cannot pre-build the full hybrid service.
// The construction cost is negligible (struct wrapping only, no I/O).
let rpc_service = NodeServiceServer::with_interceptor(make_server(), check_auth);
// Wrap S3 service with Swift service to handle Swift API requests
// Swift API is only available when compiled with the 'swift' feature
// When enabled, Swift routes are handled at /v1/AUTH_* paths by default
#[cfg(feature = "swift")]
let http_service = SwiftService::new(true, None, s3_service);
#[cfg(not(feature = "swift"))]
@@ -607,6 +588,27 @@ fn process_connection(
None
}
};
// ── Canonical Middleware Stack Order (outermost → innermost) ──
// This order MUST be preserved across refactorings.
// Only AddExtensionLayer (layers 1-2) are per-connection; layers 3-15 are stateless.
//
// 1. AddExtensionLayer<RemoteAddr> — per-connection peer address
// 2. AddExtensionLayer<SocketAddr> — per-connection raw socket addr (TrustedProxy)
// 3. TrustedProxyLayer — conditional, parses X-Forwarded-For
// 4. SetRequestIdLayer — generates X-Request-ID
// 5. RequestContextLayer — creates RequestContext in extensions
// 6. AdminChunkedContentLengthCompatLayer — admin API compat
// 7. CatchPanicLayer — panic → 500
// 8. ReadinessGateLayer — blocks until ready
// 9. KeystoneAuthLayer — X-Auth-Token validation
// 10. TraceLayer — request/response tracing + metrics
// 11. PropagateRequestIdLayer — X-Request-ID → response
// 12. PathCategoryInjectionLayer — injects path category for compression
// 13. CompressionLayer — response compression (whitelist, path-aware)
// 14. ObjectAttributesEtagFixLayer — ETag fix for GetObjectAttributes
// 15. ConditionalCorsLayer — S3 API CORS
// 16. RedirectLayer — console redirect (conditional)
// ─────────────────────────────────────────────────────────────
let hybrid_service = ServiceBuilder::new()
// NOTE: Both extension types are intentionally inserted to maintain compatibility:
// 1. `Option<RemoteAddr>` - Used by existing admin/storage handlers throughout the codebase
@@ -618,12 +620,10 @@ fn process_connection(
.option_layer(remote_addr.map(|ra| AddExtensionLayer::new(ra.0)))
// Add TrustedProxyLayer to handle X-Forwarded-For and other proxy headers
// This should be placed before TraceLayer so that logs reflect the real client IP
.option_layer(if rustfs_trusted_proxies::is_enabled() {
Some(rustfs_trusted_proxies::layer().clone())
} else {
None
})
// Pre-computed in ConnectionContext to avoid per-connection is_enabled() check.
.option_layer(trusted_proxy_layer)
.layer(SetRequestIdLayer::x_request_id(MakeRequestUuid))
.layer(RequestContextLayer)
.layer(AdminChunkedContentLengthCompatLayer)
.layer(CatchPanicLayer::new())
// CRITICAL: Insert ReadinessGateLayer before business logic
@@ -632,10 +632,8 @@ fn process_connection(
// Add Keystone authentication middleware
// This validates X-Auth-Token headers and stores credentials in task-local storage
// Must be placed AFTER ReadinessGateLayer but BEFORE business logic
.layer({
let keystone_auth = auth_keystone::get_keystone_auth();
KeystoneAuthLayer::new(keystone_auth)
})
// Pre-computed in ConnectionContext to avoid per-connection OnceLock read.
.layer(KeystoneAuthLayer::new(keystone_auth))
.layer(
TraceLayer::new_for_http()
.make_span_with(|request: &HttpRequest<_>| {
@@ -694,6 +692,12 @@ fn process_connection(
debug!("http started method: {}, url path: {}", request.method(), request.uri().path());
let labels = [("key_request_method", request.method().to_string())];
counter!("rustfs.api.requests.total", &labels).increment(1);
// Aggregate request body size for throughput monitoring (lightweight)
if let Some(cl) = request.headers().get("content-length")
&& let Some(len) = cl.to_str().ok().and_then(|s| s.parse::<u64>().ok())
{
counter!("rustfs.request.body.bytes_total", "direction" => "request").increment(len);
}
})
.on_response(|response: &Response<_>, latency: Duration, span: &Span| {
span.record("status_code", tracing::field::display(response.status()));
@@ -702,13 +706,29 @@ fn process_connection(
debug!("http response generated in {:?}", latency)
})
.on_body_chunk(|chunk: &Bytes, latency: Duration, span: &Span| {
let _enter = span.enter();
histogram!("rustfs.request.body.len").record(chunk.len() as f64);
debug!("http body sending {} bytes in {:?}", chunk.len(), latency);
// Always track aggregate body bytes (lightweight counter, no debug logging)
counter!("rustfs.request.body.bytes_total", "direction" => "response").increment(chunk.len() as u64);
#[cfg(feature = "tracing-chunk-debug")]
{
let _enter = span.enter();
histogram!("rustfs.request.body.len").record(chunk.len() as f64);
debug!("http body sending {} bytes in {:?}", chunk.len(), latency);
}
#[cfg(not(feature = "tracing-chunk-debug"))]
{
let _ = (latency, span);
}
})
.on_eos(|_trailers: Option<&HeaderMap>, stream_duration: Duration, span: &Span| {
let _enter = span.enter();
debug!("http stream closed after {:?}", stream_duration)
#[cfg(feature = "tracing-chunk-debug")]
{
let _enter = span.enter();
debug!("http stream closed after {:?}", stream_duration);
}
#[cfg(not(feature = "tracing-chunk-debug"))]
{
let _ = (_trailers, stream_duration, span);
}
})
.on_failure(|_error, latency: Duration, span: &Span| {
let _enter = span.enter();
@@ -719,7 +739,8 @@ fn process_connection(
.layer(PropagateRequestIdLayer::x_request_id())
// Compress responses based on whitelist configuration
// Only compresses when enabled and matches configured extensions/MIME types
.layer(CompressionLayer::new().compress_when(CompressionPredicate::new(compression_config)))
.layer(PathCategoryInjectionLayer)
.layer(CompressionLayer::new().compress_when(PathAwareCompressionPredicate::new(compression_config)))
.layer(ObjectAttributesEtagFixLayer)
// Conditional CORS layer: only applies to S3 API requests (not Admin, not Console)
// Admin has its own CORS handling in router.rs
@@ -733,12 +754,13 @@ fn process_connection(
let hybrid_service = TowerToHyperService::new(hybrid_service);
// Decide whether to handle HTTPS or HTTP connections based on the existence of TLS Acceptor
if let Some(acceptor) = tls_acceptor {
if let Some(holder) = tls_acceptor {
debug!("TLS handshake start");
let peer_addr = socket
.peer_addr()
.ok()
.map_or_else(|| "unknown".to_string(), |addr| addr.to_string());
let acceptor = holder.get();
match acceptor.accept(socket).await {
Ok(tls_socket) => {
debug!("TLS handshake successful");
@@ -749,32 +771,26 @@ fn process_connection(
}
}
Err(err) => {
// Detailed analysis of the reasons why the TLS handshake fails
let err_str = err.to_string();
let mut key_failure_type_str: &str = "UNKNOWN";
if err_str.contains("unexpected EOF") || err_str.contains("handshake eof") {
warn!(peer_addr = %peer_addr, "TLS handshake failed. If this client needs HTTP, it should connect to the HTTP port instead");
key_failure_type_str = "UNEXPECTED_EOF";
} else if err_str.contains("protocol version") {
error!(
peer_addr = %peer_addr,
"TLS handshake failed due to protocol version mismatch: {}", err
);
key_failure_type_str = "PROTOCOL_VERSION";
} else if err_str.contains("certificate") {
error!(
peer_addr = %peer_addr,
"TLS handshake failed due to certificate issues: {}", err
);
key_failure_type_str = "CERTIFICATE";
} else {
error!(
peer_addr = %peer_addr,
"TLS handshake failed: {}", err
);
let kind = TlsHandshakeFailureKind::classify(&err_str);
match kind {
TlsHandshakeFailureKind::UnexpectedEof => {
warn!(peer_addr = %peer_addr, "TLS handshake failed (unexpected EOF). If this client needs HTTP, it should connect to the HTTP port instead");
}
TlsHandshakeFailureKind::ProtocolVersion => {
error!(peer_addr = %peer_addr, "TLS handshake failed (protocol version mismatch): {}", err);
}
TlsHandshakeFailureKind::Certificate => {
error!(peer_addr = %peer_addr, "TLS handshake failed (certificate issue): {}", err);
}
TlsHandshakeFailureKind::Alert => {
error!(peer_addr = %peer_addr, "TLS handshake failed (alert): {}", err);
}
TlsHandshakeFailureKind::Unknown => {
error!(peer_addr = %peer_addr, "TLS handshake failed: {}", err);
}
}
counter!("rustfs_tls_handshake_failures", &[("key_failure_type", key_failure_type_str)]).increment(1);
// Record detailed diagnostic information
counter!("rustfs_tls_handshake_failures", &[("failure_type", kind.as_str())]).increment(1);
debug!(
peer_addr = %peer_addr,
error_type = %std::any::type_name_of_val(&err),
@@ -914,6 +930,68 @@ mod tests {
use http::HeaderMap;
use opentelemetry::propagation::Extractor;
/// Baseline constants — reference the authoritative config defaults.
/// If a config default changes, tests automatically follow.
mod baseline {
use rustfs_config::{
DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE, DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE, DEFAULT_H2_MAX_FRAME_SIZE,
DEFAULT_H2_MAX_HEADER_LIST_SIZE, DEFAULT_HTTP1_HEADER_READ_TIMEOUT, DEFAULT_HTTP1_MAX_BUF_SIZE,
};
/// Number of middleware layers in the canonical stack order (see http.rs).
/// Layers 1-2 are per-connection (AddExtension), 3-15 are stateless.
pub const MIDDLEWARE_LAYER_COUNT: usize = 15;
/// Current HTTP/2 defaults (from rustfs_config).
pub const H2_INITIAL_STREAM_WINDOW_SIZE: u32 = DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE;
pub const H2_INITIAL_CONN_WINDOW_SIZE: u32 = DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE;
pub const H2_MAX_FRAME_SIZE: u32 = DEFAULT_H2_MAX_FRAME_SIZE;
pub const H2_MAX_HEADER_LIST_SIZE: u32 = DEFAULT_H2_MAX_HEADER_LIST_SIZE;
/// Current HTTP/1.1 defaults (from rustfs_config).
pub const HTTP1_HEADER_READ_TIMEOUT_SECS: u64 = DEFAULT_HTTP1_HEADER_READ_TIMEOUT;
pub const HTTP1_MAX_BUF_SIZE: usize = DEFAULT_HTTP1_MAX_BUF_SIZE;
/// Post-accept socket syscalls after T03 optimization.
/// Linux: 1 (TCP_QUICKACK only). Other platforms: 0.
#[cfg(target_os = "linux")]
pub const POST_ACCEPT_SYSCALL_COUNT_LINUX: usize = 1;
#[cfg(not(target_os = "linux"))]
pub const POST_ACCEPT_SYSCALL_COUNT_OTHER: usize = 0;
}
#[test]
fn test_baseline_h2_constants() {
use rustfs_config::{
DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE, DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE, DEFAULT_H2_MAX_FRAME_SIZE,
DEFAULT_H2_MAX_HEADER_LIST_SIZE,
};
assert_eq!(baseline::H2_INITIAL_STREAM_WINDOW_SIZE, DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE);
assert_eq!(baseline::H2_INITIAL_CONN_WINDOW_SIZE, DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE);
assert_eq!(baseline::H2_MAX_FRAME_SIZE, DEFAULT_H2_MAX_FRAME_SIZE);
assert_eq!(baseline::H2_MAX_HEADER_LIST_SIZE, DEFAULT_H2_MAX_HEADER_LIST_SIZE);
}
#[test]
fn test_baseline_http1_constants() {
use rustfs_config::{DEFAULT_HTTP1_HEADER_READ_TIMEOUT, DEFAULT_HTTP1_MAX_BUF_SIZE};
assert_eq!(baseline::HTTP1_HEADER_READ_TIMEOUT_SECS, DEFAULT_HTTP1_HEADER_READ_TIMEOUT);
assert_eq!(baseline::HTTP1_MAX_BUF_SIZE, DEFAULT_HTTP1_MAX_BUF_SIZE);
}
#[test]
fn test_baseline_middleware_count() {
assert_eq!(baseline::MIDDLEWARE_LAYER_COUNT, 15);
}
#[test]
fn test_baseline_post_accept_syscall_count() {
#[cfg(target_os = "linux")]
assert_eq!(baseline::POST_ACCEPT_SYSCALL_COUNT_LINUX, 1);
#[cfg(not(target_os = "linux"))]
assert_eq!(baseline::POST_ACCEPT_SYSCALL_COUNT_OTHER, 0);
}
#[test]
fn test_headermap_carrier_new() {
let headers = HeaderMap::new();
+177
View File
@@ -17,19 +17,134 @@ use crate::server::cors;
use crate::server::hybrid::HybridBody;
use crate::server::{ADMIN_PREFIX, CONSOLE_PREFIX, MINIO_ADMIN_PREFIX, MINIO_ADMIN_V3_PREFIX, RPC_PREFIX, RUSTFS_ADMIN_PREFIX};
use crate::storage::apply_cors_headers;
use crate::storage::request_context::{RequestContext, extract_request_id_from_headers};
use bytes::Bytes;
use http::{HeaderMap, HeaderValue, Method, Request as HttpRequest, Response, StatusCode};
use http_body::Body;
use http_body_util::BodyExt;
use hyper::body::Incoming;
use opentelemetry::global;
use opentelemetry::trace::TraceContextExt;
use rustfs_utils::get_env_opt_str;
use rustfs_utils::http::headers::AMZ_REQUEST_ID;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Instant;
use tower::{Layer, Service};
use tracing::debug;
/// A carrier that adapts [`HeaderMap`] for OpenTelemetry trace context propagation.
struct HeaderMapCarrier<'a>(&'a HeaderMap);
impl<'a> opentelemetry::propagation::Extractor for HeaderMapCarrier<'a> {
fn get(&self, key: &str) -> Option<&str> {
self.0.get(key).and_then(|v| v.to_str().ok())
}
fn keys(&self) -> Vec<&str> {
self.0.keys().map(|k| k.as_str()).collect()
}
fn get_all(&self, key: &str) -> Option<Vec<&str>> {
let headers = self
.0
.get_all(key)
.iter()
.filter_map(|value| value.to_str().ok())
.collect::<Vec<_>>();
if headers.is_empty() { None } else { Some(headers) }
}
}
/// Tower middleware layer that creates a canonical [`RequestContext`] from HTTP headers
/// and injects it into `request.extensions()`.
///
/// This layer must be placed after `SetRequestIdLayer` in the middleware stack,
/// as it reads the `x-request-id` header that `SetRequestIdLayer` generates.
///
/// Additionally, it sets the `x-amz-request-id` request header for S3 compatibility
/// if not already present.
#[derive(Clone, Default)]
pub struct RequestContextLayer;
impl<S> Layer<S> for RequestContextLayer {
type Service = RequestContextService<S>;
fn layer(&self, inner: S) -> Self::Service {
RequestContextService { inner }
}
}
/// Service that injects [`RequestContext`] into every request.
#[derive(Clone)]
pub struct RequestContextService<S> {
inner: S,
}
impl<S, B> Service<HttpRequest<B>> for RequestContextService<S>
where
S: Service<HttpRequest<B>>,
{
type Response = S::Response;
type Error = S::Error;
type Future = S::Future;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: HttpRequest<B>) -> Self::Future {
let request_id = extract_request_id_from_headers(req.headers());
// Extract OpenTelemetry trace/span context from incoming headers
let parent_cx = global::get_text_map_propagator(|propagator| propagator.extract(&HeaderMapCarrier(req.headers())));
let span_ref = parent_cx.span();
let span_context = span_ref.span_context();
let trace_id = if span_context.is_valid() {
Some(span_context.trace_id().to_string())
} else {
None
};
let span_id = if span_context.is_valid() {
Some(span_context.span_id().to_string())
} else {
None
};
// Preserve the upstream x-amz-request-id if present (S3 client forwarding),
// otherwise fall back to the canonical request_id.
let x_amz_request_id = req
.headers()
.get(AMZ_REQUEST_ID)
.and_then(|v| v.to_str().ok())
.map(String::from)
.unwrap_or_else(|| request_id.clone());
let ctx = RequestContext {
request_id: request_id.clone(),
x_amz_request_id,
trace_id,
span_id,
start_time: Instant::now(),
};
req.extensions_mut().insert(ctx);
// Set x-amz-request-id for S3 compatibility downstream
if !req.headers().contains_key(AMZ_REQUEST_ID)
&& let Ok(val) = HeaderValue::from_str(&request_id)
{
req.headers_mut()
.insert(http::header::HeaderName::from_static(AMZ_REQUEST_ID), val);
}
self.inner.call(req)
}
}
/// Redirect layer that redirects browser requests to the console
#[derive(Clone)]
pub struct RedirectLayer;
@@ -563,11 +678,30 @@ where
#[cfg(test)]
mod tests {
use super::*;
use futures::future::{Ready, ready};
use http::Request;
use http_body_util::BodyExt;
use http_body_util::Full;
use std::convert::Infallible;
use temp_env::with_var;
#[derive(Clone, Debug)]
struct CaptureService;
impl<B> Service<Request<B>> for CaptureService {
type Response = Request<B>;
type Error = Infallible;
type Future = Ready<Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: Request<B>) -> Self::Future {
ready(Ok(req))
}
}
#[test]
fn admin_chunked_put_without_content_length_is_normalized() {
let request = Request::builder()
@@ -701,6 +835,49 @@ mod tests {
});
}
#[test]
fn request_context_layer_populates_context_and_s3_request_id_from_x_request_id() {
let mut service = RequestContextLayer.layer(CaptureService);
let request = Request::builder()
.uri("/bucket/object")
.header("x-request-id", "req-123")
.body(())
.expect("request");
let request = service.call(request).into_inner().expect("service call should succeed");
let context = request
.extensions()
.get::<RequestContext>()
.expect("request context should be present");
assert_eq!(context.request_id, "req-123");
assert_eq!(context.x_amz_request_id, "req-123");
assert!(context.trace_id.is_none());
assert!(context.span_id.is_none());
assert_eq!(request.headers().get(AMZ_REQUEST_ID).unwrap(), "req-123");
}
#[test]
fn request_context_layer_preserves_upstream_s3_request_id() {
let mut service = RequestContextLayer.layer(CaptureService);
let request = Request::builder()
.uri("/bucket/object")
.header("x-request-id", "req-123")
.header(AMZ_REQUEST_ID, "amz-456")
.body(())
.expect("request");
let request = service.call(request).into_inner().expect("service call should succeed");
let context = request
.extensions()
.get::<RequestContext>()
.expect("request context should be present");
assert_eq!(context.request_id, "req-123");
assert_eq!(context.x_amz_request_id, "amz-456");
assert_eq!(request.headers().get(AMZ_REQUEST_ID).unwrap(), "amz-456");
}
#[tokio::test]
async fn test_resolve_s3_options_cors_headers_no_headers_without_match() {
let mut req_headers = HeaderMap::new();
+1 -2
View File
@@ -13,7 +13,6 @@
// limitations under the License.
mod audit;
mod cert;
mod compress;
pub mod cors;
mod event;
@@ -24,9 +23,9 @@ mod prefix;
mod readiness;
mod runtime;
mod service_state;
pub(crate) mod tls_material;
pub(crate) use audit::{start_audit_system, stop_audit_system};
pub(crate) use cert::init_cert;
pub(crate) use event::{init_event_notifier, shutdown_event_notifier};
pub(crate) use http::start_http_server;
pub(crate) use prefix::*;
+510
View File
@@ -0,0 +1,510 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! Unified TLS Material Snapshot
//!
//! Provides a single loading point for all TLS materials, eliminating duplicate
//! directory scanning and PEM parsing between outbound and inbound paths.
//!
//! Usage:
//! 1. Call `TlsMaterialSnapshot::load(tls_path)` once at startup.
//! 2. Call `snapshot.apply_outbound()` to set global root CAs and mTLS identity.
//! 3. Call `snapshot.build_tls_acceptor(tls_path)` to build the server TLS acceptor.
use rustfs_common::{MtlsIdentityPem, set_global_mtls_identity, set_global_root_cert};
use rustfs_config::{
DEFAULT_TLS_RELOAD_ENABLE, DEFAULT_TLS_RELOAD_INTERVAL, DEFAULT_TRUST_LEAF_CERT_AS_CA, DEFAULT_TRUST_SYSTEM_CA,
ENV_MTLS_CLIENT_CERT, ENV_MTLS_CLIENT_KEY, ENV_TLS_RELOAD_ENABLE, ENV_TLS_RELOAD_INTERVAL, ENV_TRUST_LEAF_CERT_AS_CA,
ENV_TRUST_SYSTEM_CA, RUSTFS_CA_CERT, RUSTFS_CLIENT_CERT_FILENAME, RUSTFS_CLIENT_KEY_FILENAME, RUSTFS_PUBLIC_CERT,
RUSTFS_TLS_CERT, RUSTFS_TLS_KEY,
};
use rustfs_utils::{get_env_bool, get_env_opt_str};
use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject};
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::RwLock;
use std::time::Duration;
use tokio_rustls::TlsAcceptor;
use tracing::{debug, info, warn};
/// System CA certificate search paths (platform-specific).
const SYSTEM_CA_PATHS: &[&str] = &[
"/etc/ssl/certs/ca-certificates.crt", // Debian/Ubuntu/Alpine
"/etc/pki/tls/certs/ca-bundle.crt", // Fedora/RHEL/CentOS
"/etc/ssl/ca-bundle.pem", // OpenSUSE
"/etc/pki/tls/cacert.pem", // OpenELEC
"/etc/ssl/cert.pem", // macOS/FreeBSD
"/usr/local/etc/openssl/cert.pem", // macOS/Homebrew OpenSSL
"/usr/local/share/certs/ca-root-nss.crt", // FreeBSD
"/etc/pki/ca-trust/extracted/pem/tls-ca-bundle.pem", // RHEL
"/usr/share/pki/ca-trust-legacy/ca-bundle.legacy.crt", // RHEL legacy
];
/// Outbound TLS material for client connections (inter-node RPC).
#[derive(Debug, Clone)]
pub struct OutboundTlsMaterial {
/// Concatenated PEM-encoded root CA certificates.
pub root_ca_pem: Vec<u8>,
/// Optional mTLS client identity.
pub mtls_identity: Option<MtlsIdentityPem>,
}
/// Complete TLS material snapshot loaded once at startup.
#[derive(Debug)]
pub struct TlsMaterialSnapshot {
/// Material for outbound client connections.
pub outbound: OutboundTlsMaterial,
/// Whether any server certificates were found.
pub has_server_certs: bool,
}
impl TlsMaterialSnapshot {
/// Load all TLS materials from the given directory.
///
/// This is the single entry point that replaces both the old
/// `cert.rs::init_cert()` and `http.rs::setup_tls_acceptor()` loading logic.
pub async fn load(tls_path: &str) -> Result<Self, TlsMaterialError> {
if tls_path.is_empty() {
info!("No TLS path configured; skipping TLS material loading");
return Ok(Self::empty());
}
let tls_dir = PathBuf::from(tls_path);
// Load outbound material (root CAs + mTLS identity)
let outbound = load_outbound_material(&tls_dir).await?;
// Check if server certs exist (actual loading happens in build_tls_acceptor)
let has_server_certs = has_server_certificates(tls_path).await;
Ok(Self {
outbound,
has_server_certs,
})
}
/// Apply outbound material to global state (root CAs, mTLS identity).
pub async fn apply_outbound(&self) {
if !self.outbound.root_ca_pem.is_empty() {
set_global_root_cert(self.outbound.root_ca_pem.clone()).await;
info!("Configured custom root certificates for inter-node communication");
}
set_global_mtls_identity(self.outbound.mtls_identity.clone()).await;
}
/// Build a `TlsAcceptorHolder` from the loaded snapshot.
///
/// This is the single place that constructs the server `ServerConfig`,
/// handling both multi-cert (SNI resolver) and single-cert fallback.
/// Returns `None` if no TLS certificates are available.
pub async fn build_tls_acceptor(&self, tls_path: &str) -> Result<Option<Arc<TlsAcceptorHolder>>, TlsMaterialError> {
if tls_path.is_empty() || !self.has_server_certs {
return Ok(None);
}
let mtls_verifier = rustfs_utils::build_webpki_client_verifier(tls_path)
.map_err(|e| TlsMaterialError::Io(format!("build mTLS verifier: {e}")))?;
// Try multi-cert (SNI) first
match rustfs_utils::load_all_certs_from_directory(tls_path) {
Ok(cert_key_pairs) if !cert_key_pairs.is_empty() => match rustfs_utils::create_multi_cert_resolver(cert_key_pairs) {
Ok(resolver) => {
let config = build_server_config(ServerCertSource::Resolver(Arc::new(resolver)), mtls_verifier)?;
info!("Created TLS acceptor with SNI resolver");
let acceptor = Arc::new(TlsAcceptor::from(Arc::new(config)));
return Ok(Some(Arc::new(TlsAcceptorHolder::new(acceptor))));
}
Err(e) => warn!("Failed to build multi-cert resolver: {}, falling back to single-cert", e),
},
Ok(_) => debug!("No valid multi-cert directory structure found"),
Err(_) => debug!("load_all_certs_from_directory failed, trying single-cert fallback"),
}
// Fallback: single cert
let key_path = format!("{tls_path}/{RUSTFS_TLS_KEY}");
let cert_path = format!("{tls_path}/{RUSTFS_TLS_CERT}");
if tokio::try_join!(tokio::fs::metadata(&key_path), tokio::fs::metadata(&cert_path)).is_ok() {
let certs = rustfs_utils::load_certs(&cert_path).map_err(|e| TlsMaterialError::Io(format!("load certs: {e}")))?;
let key = rustfs_utils::load_private_key(&key_path).map_err(|e| TlsMaterialError::Io(format!("load key: {e}")))?;
let config = build_server_config(ServerCertSource::SingleCert { certs, key }, mtls_verifier)?;
info!("Created TLS acceptor with single certificate");
let acceptor = Arc::new(TlsAcceptor::from(Arc::new(config)));
return Ok(Some(Arc::new(TlsAcceptorHolder::new(acceptor))));
}
debug!("No valid TLS certificates found, starting with HTTP");
Ok(None)
}
fn empty() -> Self {
Self {
outbound: OutboundTlsMaterial {
root_ca_pem: Vec::new(),
mtls_identity: None,
},
has_server_certs: false,
}
}
}
// ── Server Config Construction ──
/// Certificate source for building a `ServerConfig`.
enum ServerCertSource {
/// Pre-built SNI resolver from multi-cert directory.
Resolver(Arc<dyn rustls::server::ResolvesServerCert + Send + Sync>),
/// Single certificate/key pair.
SingleCert {
certs: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
},
}
/// Build a `ServerConfig` with standardized ALPN, session cache, and key log settings.
///
/// This is the single place for `ServerConfig` construction, used by both
/// initial startup and hot-reload.
fn build_server_config(
cert_source: ServerCertSource,
mtls_verifier: Option<Arc<dyn rustls::server::danger::ClientCertVerifier>>,
) -> Result<rustls::ServerConfig, TlsMaterialError> {
let mut config = match cert_source {
ServerCertSource::Resolver(resolver) => {
if let Some(verifier) = mtls_verifier {
rustls::ServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_cert_resolver(resolver)
} else {
rustls::ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(resolver)
}
}
ServerCertSource::SingleCert { certs, key } => {
if let Some(verifier) = mtls_verifier {
rustls::ServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_single_cert(certs, key)
.map_err(|e| TlsMaterialError::Io(format!("configure single cert with mTLS: {e}")))?
} else {
rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.map_err(|e| TlsMaterialError::Io(format!("configure single cert: {e}")))?
}
}
};
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec(), b"http/1.0".to_vec()];
config.session_storage = rustls::server::ServerSessionMemoryCache::new(10000);
if rustfs_utils::tls_key_log() {
config.key_log = Arc::new(rustls::KeyLogFile::new());
}
Ok(config)
}
// ── Outbound Material Loading ──
/// Load root CA certificates and mTLS identity for outbound connections.
async fn load_outbound_material(tls_dir: &Path) -> Result<OutboundTlsMaterial, TlsMaterialError> {
let mut root_ca_pem = Vec::new();
// 1. Optional: load leaf certs as root CAs
if get_env_bool(ENV_TRUST_LEAF_CERT_AS_CA, DEFAULT_TRUST_LEAF_CERT_AS_CA)
&& load_cert_file_by_name(tls_dir, RUSTFS_TLS_CERT, &mut root_ca_pem).await
{
info!("Loaded leaf certificate(s) as root CA as per RUSTFS_TRUST_LEAF_CERT_AS_CA");
}
// 2. Load public.crt and ca.crt
load_cert_file(&tls_dir.join(RUSTFS_PUBLIC_CERT), &mut root_ca_pem, "CA certificate").await;
load_cert_file(&tls_dir.join(RUSTFS_CA_CERT), &mut root_ca_pem, "CA certificate").await;
// 3. Optional: load system root CAs
if get_env_bool(ENV_TRUST_SYSTEM_CA, DEFAULT_TRUST_SYSTEM_CA) {
let mut system_loaded = false;
for path in SYSTEM_CA_PATHS {
if load_cert_file(Path::new(path), &mut root_ca_pem, "system root certificates").await {
system_loaded = true;
info!("Loaded system root certificates from {}", path);
break;
}
}
if !system_loaded {
debug!("Could not find system root certificates in common locations.");
}
} else {
info!("Loading system root certificates disabled via RUSTFS_TRUST_SYSTEM_CA");
}
// 4. Load optional mTLS identity
let mtls_identity = load_mtls_identity(tls_dir).await?;
Ok(OutboundTlsMaterial {
root_ca_pem,
mtls_identity,
})
}
/// Quick check whether server certificate files exist in the TLS directory.
async fn has_server_certificates(tls_path: &str) -> bool {
if tokio::fs::metadata(tls_path).await.is_err() {
return false;
}
// Check for multi-cert directory structure OR single cert files
if rustfs_utils::load_all_certs_from_directory(tls_path).is_ok_and(|p| !p.is_empty()) {
return true;
}
let key_path = format!("{tls_path}/{RUSTFS_TLS_KEY}");
let cert_path = format!("{tls_path}/{RUSTFS_TLS_CERT}");
tokio::try_join!(tokio::fs::metadata(&key_path), tokio::fs::metadata(&cert_path)).is_ok()
}
/// Load mTLS client identity from the TLS directory.
async fn load_mtls_identity(tls_dir: &Path) -> Result<Option<MtlsIdentityPem>, TlsMaterialError> {
let client_cert_path = match get_env_opt_str(ENV_MTLS_CLIENT_CERT) {
Some(p) => PathBuf::from(p),
None => tls_dir.join(RUSTFS_CLIENT_CERT_FILENAME),
};
let client_key_path = match get_env_opt_str(ENV_MTLS_CLIENT_KEY) {
Some(p) => PathBuf::from(p),
None => tls_dir.join(RUSTFS_CLIENT_KEY_FILENAME),
};
if !client_cert_path.exists() || !client_key_path.exists() {
info!(
"mTLS client identity not configured (missing {:?} and/or {:?}); proceeding with server-only TLS",
client_cert_path, client_key_path
);
return Ok(None);
}
let cert_pem = tokio::fs::read(&client_cert_path)
.await
.map_err(|e| TlsMaterialError::Io(format!("read client cert {client_cert_path:?}: {e}")))?;
let key_pem = tokio::fs::read(&client_key_path)
.await
.map_err(|e| TlsMaterialError::Io(format!("read client key {client_key_path:?}: {e}")))?;
// Validate parse-ability
let mut reader = std::io::Cursor::new(&cert_pem);
if CertificateDer::pem_reader_iter(&mut reader).next().is_none() {
return Err(TlsMaterialError::Parse("no valid certificate in client cert PEM".into()));
}
let mut reader = std::io::Cursor::new(&key_pem);
PrivateKeyDer::from_pem_reader(&mut reader).map_err(|e| TlsMaterialError::Parse(format!("invalid client key PEM: {e}")))?;
info!("Loaded mTLS client identity cert={:?} key={:?}", client_cert_path, client_key_path);
Ok(Some(MtlsIdentityPem { cert_pem, key_pem }))
}
/// Load a single certificate file and append PEM data.
/// Returns true if the file was successfully loaded.
async fn load_cert_file(path: &Path, pem_data: &mut Vec<u8>, desc: &str) -> bool {
if tokio::fs::metadata(path).await.is_err() {
debug!("{} file not found at {:?}", desc, path);
return false;
}
match tokio::fs::read(path).await {
Ok(data) => {
pem_data.extend_from_slice(&data);
pem_data.push(b'\n');
info!("Loaded {} from {:?}", desc, path);
true
}
Err(e) => {
debug!("Failed to read {} from {:?}: {}", desc, path, e);
false
}
}
}
/// Search for and load certificate files matching `cert_name` in the directory
/// and one level of subdirectories.
/// Returns `true` if at least one matching file was loaded.
async fn load_cert_file_by_name(dir: &Path, cert_name: &str, pem_data: &mut Vec<u8>) -> bool {
let Ok(mut rd) = tokio::fs::read_dir(dir).await else {
debug!("Certificate directory not found: {}", dir.display());
return false;
};
let mut loaded = false;
while let Ok(Some(entry)) = rd.next_entry().await {
let Ok(ft) = entry.file_type().await else { continue };
if ft.is_file() {
let fname = entry.file_name().to_string_lossy().to_string();
if fname == cert_name && load_cert_file(&entry.path(), pem_data, "certificate").await {
loaded = true;
}
} else if ft.is_dir() {
// Only check direct subdirectories (one level deep)
if let Ok(mut sub_rd) = tokio::fs::read_dir(&entry.path()).await {
while let Ok(Some(sub_entry)) = sub_rd.next_entry().await {
if let Ok(sub_ft) = sub_entry.file_type().await
&& sub_ft.is_file()
{
let fname = sub_entry.file_name().to_string_lossy().to_string();
if fname == cert_name && load_cert_file(&sub_entry.path(), pem_data, "certificate").await {
loaded = true;
}
}
}
}
}
}
loaded
}
/// Errors that can occur during TLS material loading.
#[derive(Debug)]
pub enum TlsMaterialError {
/// I/O error (file read, directory access).
Io(String),
/// PEM parsing error.
Parse(String),
}
impl std::fmt::Display for TlsMaterialError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TlsMaterialError::Io(msg) => write!(f, "TLS material I/O error: {msg}"),
TlsMaterialError::Parse(msg) => write!(f, "TLS material parse error: {msg}"),
}
}
}
impl std::error::Error for TlsMaterialError {}
// ── TLS Handshake Error Classification ──
/// Structured classification of TLS handshake failures.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum TlsHandshakeFailureKind {
UnexpectedEof,
ProtocolVersion,
Certificate,
Alert,
Unknown,
}
impl TlsHandshakeFailureKind {
/// Classify a TLS accept error into a structured failure kind.
pub(crate) fn classify(err_msg: &str) -> Self {
if err_msg.contains("unexpected EOF") || err_msg.contains("handshake eof") {
Self::UnexpectedEof
} else if err_msg.contains("protocol version") {
Self::ProtocolVersion
} else if err_msg.contains("certificate") || err_msg.contains("invalid peer certificate") {
Self::Certificate
} else if err_msg.contains("alert") {
Self::Alert
} else {
Self::Unknown
}
}
/// Metric label string for Prometheus.
pub(crate) fn as_str(self) -> &'static str {
match self {
Self::UnexpectedEof => "UNEXPECTED_EOF",
Self::ProtocolVersion => "PROTOCOL_VERSION",
Self::Certificate => "CERTIFICATE",
Self::Alert => "ALERT",
Self::Unknown => "UNKNOWN",
}
}
}
// ── TLS Acceptor Holder (for hot reload) ──
/// Holds the current TLS acceptor and supports atomic swap for certificate rotation.
///
/// Uses `RwLock` so that multiple readers (per-connection `get()` calls)
/// do not block each other. The write lock is held only briefly during swap.
pub(crate) struct TlsAcceptorHolder {
current: RwLock<Arc<TlsAcceptor>>,
}
impl TlsAcceptorHolder {
pub(crate) fn new(acceptor: Arc<TlsAcceptor>) -> Self {
Self {
current: RwLock::new(acceptor),
}
}
/// Get the current TLS acceptor for handling a new connection.
#[inline]
pub(crate) fn get(&self) -> Arc<TlsAcceptor> {
match self.current.read() {
Ok(guard) => guard.clone(),
Err(poisoned) => poisoned.into_inner().clone(),
}
}
/// Atomically replace the TLS acceptor with a new one.
fn swap(&self, new_holder: &TlsAcceptorHolder) {
let new_acceptor = new_holder.get();
match self.current.write() {
Ok(mut guard) => *guard = new_acceptor,
Err(poisoned) => {
let mut guard = poisoned.into_inner();
*guard = new_acceptor;
}
}
}
}
/// Spawn a background task that periodically checks for TLS certificate changes.
pub(crate) fn spawn_reload_loop(tls_path: String, holder: Arc<TlsAcceptorHolder>) {
let enabled = get_env_bool(ENV_TLS_RELOAD_ENABLE, DEFAULT_TLS_RELOAD_ENABLE);
if !enabled {
debug!("TLS certificate hot reload is disabled (set {}=1 to enable)", ENV_TLS_RELOAD_ENABLE);
return;
}
let interval_secs = rustfs_utils::get_env_u64(ENV_TLS_RELOAD_INTERVAL, DEFAULT_TLS_RELOAD_INTERVAL).max(5);
info!("TLS certificate hot reload enabled, checking every {}s", interval_secs);
tokio::spawn(async move {
let mut interval = tokio::time::interval(Duration::from_secs(interval_secs));
loop {
interval.tick().await;
match TlsMaterialSnapshot::load(&tls_path).await {
Ok(snapshot) => {
// Always refresh outbound material (root CAs, mTLS identity) on reload.
snapshot.apply_outbound().await;
match snapshot.build_tls_acceptor(&tls_path).await {
Ok(Some(new_holder)) => {
info!("TLS certificates reloaded successfully");
holder.swap(&new_holder);
}
Ok(None) => debug!("TLS reload: no server certificates found in directory, skipping"),
Err(e) => warn!("TLS certificate reload failed (will retry): {}", e),
}
}
Err(e) => {
warn!("TLS material reload failed (will retry): {}", e);
}
}
}
});
}
+14
View File
@@ -17,6 +17,7 @@ use crate::auth::{check_key_valid, get_condition_values_with_query, get_session_
use crate::error::ApiError;
use crate::license::license_check;
use crate::server::RemoteAddr;
use crate::storage::request_context::RequestContext;
use metrics::counter;
use rustfs_ecstore::bucket::metadata_sys;
use rustfs_ecstore::bucket::policy_sys::PolicySys;
@@ -45,6 +46,7 @@ pub(crate) struct ReqInfo {
pub version_id: Option<String>,
#[allow(dead_code)]
pub region: Option<s3s::region::Region>,
pub request_context: Option<RequestContext>,
}
#[derive(Clone, Debug)]
@@ -67,6 +69,15 @@ fn ext_req_info_mut(ext: &mut http::Extensions) -> S3Result<&mut ReqInfo> {
.ok_or_else(|| s3_error!(InternalError, "ReqInfo not found in request extensions"))
}
/// Extract the canonical `RequestContext` from a request, checking both
/// the request extensions directly and the `ReqInfo.request_context` field.
pub(crate) fn request_context_from_req<T>(req: &S3Request<T>) -> Option<RequestContext> {
req.extensions
.get::<RequestContext>()
.cloned()
.or_else(|| req.extensions.get::<ReqInfo>().and_then(|ri| ri.request_context.clone()))
}
#[derive(Clone, Debug)]
pub(crate) struct ObjectTagConditions {
bucket: String,
@@ -731,10 +742,13 @@ impl S3Access for FS {
(None, false)
};
let request_context = cx.extensions_mut().get::<RequestContext>().cloned();
let req_info = ReqInfo {
cred,
is_owner,
region: rustfs_ecstore::global::get_global_region(),
request_context,
..Default::default()
};
@@ -32,6 +32,7 @@
use hashbrown::HashMap;
use moka::future::Cache;
use rustfs_config::MI_B;
use rustfs_object_io::get::GetObjectCacheWriteback;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
@@ -1063,6 +1064,13 @@ pub struct CachedGetObject {
pub replication_status: Option<String>,
/// User-defined metadata (x-amz-meta-*)
pub user_metadata: std::collections::HashMap<String, String>,
/// Additional checksum metadata persisted with cached GET responses
pub checksum_crc32: Option<String>,
pub checksum_crc32c: Option<String>,
pub checksum_sha1: Option<String>,
pub checksum_sha256: Option<String>,
pub checksum_crc64nvme: Option<String>,
pub checksum_type: Option<s3s::dto::ChecksumType>,
/// When this object was cached (for internal use, automatically set)
#[allow(dead_code)]
cached_at: Option<Instant>,
@@ -1089,6 +1097,12 @@ impl Default for CachedGetObject {
tag_count: None,
replication_status: None,
user_metadata: std::collections::HashMap::new(),
checksum_crc32: None,
checksum_crc32c: None,
checksum_sha1: None,
checksum_sha256: None,
checksum_crc64nvme: None,
checksum_type: None,
cached_at: None,
access_count: Arc::new(AtomicU64::new(0)),
}
@@ -1109,6 +1123,35 @@ impl CachedGetObject {
}
}
/// Consume a GET cache writeback payload into the cache-owned representation.
pub fn from_get_object_cache_writeback(writeback: GetObjectCacheWriteback) -> Self {
Self {
body: writeback.body,
content_length: writeback.content_length,
content_type: writeback.content_type,
e_tag: writeback.e_tag,
last_modified: writeback.last_modified,
expires: writeback.expires,
cache_control: writeback.cache_control,
content_disposition: writeback.content_disposition,
content_encoding: writeback.content_encoding,
content_language: writeback.content_language,
storage_class: writeback.storage_class,
version_id: writeback.version_id,
delete_marker: writeback.delete_marker,
user_metadata: writeback.user_metadata,
checksum_crc32: writeback.checksum_crc32,
checksum_crc32c: writeback.checksum_crc32c,
checksum_sha1: writeback.checksum_sha1,
checksum_sha256: writeback.checksum_sha256,
checksum_crc64nvme: writeback.checksum_crc64nvme,
checksum_type: writeback.checksum_type,
cached_at: Some(Instant::now()),
access_count: Arc::new(AtomicU64::new(0)),
..Default::default()
}
}
/// Builder method to set content_type
pub fn with_content_type(mut self, content_type: String) -> Self {
self.content_type = Some(content_type);
@@ -1881,6 +1924,57 @@ mod cached_object_tests {
assert_eq!(obj.user_metadata.get("x-amz-meta-custom"), Some(&"value".to_string()));
}
#[test]
fn test_cached_get_object_from_get_object_cache_writeback() {
let body = Arc::new(Bytes::from("test data"));
let obj = CachedGetObject::from_get_object_cache_writeback(GetObjectCacheWriteback {
body: Arc::clone(&body),
content_length: 9,
content_type: Some("text/plain".to_string()),
content_encoding: Some("gzip".to_string()),
cache_control: Some("max-age=3600".to_string()),
content_disposition: Some("attachment".to_string()),
content_language: Some("en-US".to_string()),
expires: Some("2024-12-31T23:59:59Z".to_string()),
storage_class: Some("STANDARD".to_string()),
version_id: Some("null".to_string()),
delete_marker: false,
user_metadata: {
let mut metadata = std::collections::HashMap::new();
metadata.insert("custom-key".to_string(), "value".to_string());
metadata
},
e_tag: Some("\"abc123\"".to_string()),
last_modified: Some("2024-01-01T12:00:00Z".to_string()),
checksum_crc32: Some("crc32".to_string()),
checksum_crc32c: None,
checksum_sha1: None,
checksum_sha256: None,
checksum_crc64nvme: None,
checksum_type: Some(s3s::dto::ChecksumType::from_static(s3s::dto::ChecksumType::FULL_OBJECT)),
});
assert_eq!(obj.content_length, 9);
assert_eq!(obj.content_type.as_deref(), Some("text/plain"));
assert_eq!(obj.content_encoding.as_deref(), Some("gzip"));
assert_eq!(obj.cache_control.as_deref(), Some("max-age=3600"));
assert_eq!(obj.content_disposition.as_deref(), Some("attachment"));
assert_eq!(obj.content_language.as_deref(), Some("en-US"));
assert_eq!(obj.expires.as_deref(), Some("2024-12-31T23:59:59Z"));
assert_eq!(obj.storage_class.as_deref(), Some("STANDARD"));
assert_eq!(obj.version_id.as_deref(), Some("null"));
assert!(!obj.delete_marker);
assert_eq!(obj.user_metadata.get("custom-key").map(String::as_str), Some("value"));
assert_eq!(obj.e_tag.as_deref(), Some("\"abc123\""));
assert_eq!(obj.last_modified.as_deref(), Some("2024-01-01T12:00:00Z"));
assert_eq!(obj.checksum_crc32.as_deref(), Some("crc32"));
assert_eq!(
obj.checksum_type,
Some(s3s::dto::ChecksumType::from_static(s3s::dto::ChecksumType::FULL_OBJECT))
);
assert!(Arc::ptr_eq(&obj.body, &body));
}
#[test]
fn test_cached_get_object_size() {
let obj = CachedGetObject::new(Bytes::from("test"), 4);
-39
View File
@@ -29,7 +29,6 @@ use rustfs_ecstore::{
use rustfs_s3_common::{S3Operation, record_s3_op};
use s3s::{S3, S3Error, S3ErrorCode, S3Request, S3Response, S3Result, dto::*, s3_error};
use std::fmt::Debug;
use tokio::io::{AsyncRead, AsyncSeek};
use tracing::{debug, error, instrument, warn};
use uuid::Uuid;
@@ -44,44 +43,6 @@ pub(crate) struct ListObjectUnorderedQuery {
pub(crate) allow_unordered: Option<String>,
}
pub(crate) struct InMemoryAsyncReader {
cursor: std::io::Cursor<Vec<u8>>,
}
impl InMemoryAsyncReader {
pub(crate) fn new(data: Vec<u8>) -> Self {
Self {
cursor: std::io::Cursor::new(data),
}
}
}
impl AsyncRead for InMemoryAsyncReader {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let unfilled = buf.initialize_unfilled();
let bytes_read = std::io::Read::read(&mut self.cursor, unfilled)?;
buf.advance(bytes_read);
std::task::Poll::Ready(Ok(()))
}
}
impl AsyncSeek for InMemoryAsyncReader {
fn start_seek(mut self: std::pin::Pin<&mut Self>, position: std::io::SeekFrom) -> std::io::Result<()> {
// std::io::Cursor natively supports negative SeekCurrent offsets
// It will automatically handle validation and return an error if the final position would be negative
std::io::Seek::seek(&mut self.cursor, position)?;
Ok(())
}
fn poll_complete(self: std::pin::Pin<&mut Self>, _cx: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<u64>> {
std::task::Poll::Ready(Ok(self.cursor.position()))
}
}
impl FS {
pub fn new() -> Self {
rustfs_s3_common::init_s3_metrics();
+124 -13
View File
@@ -12,7 +12,9 @@
// See the License for the specific language governing permissions and
// limitations under the License.
use crate::storage::access::ReqInfo;
use crate::storage::access::{ReqInfo, request_context_from_req};
use crate::storage::request_context::{RequestContext, extract_request_id_from_headers};
use hashbrown::HashMap;
use http::StatusCode;
use rustfs_audit::{
entity::{ApiDetails, ApiDetailsBuilder, AuditEntryBuilder},
@@ -24,14 +26,17 @@ use rustfs_s3_common::record_s3_op;
use rustfs_s3_common::{EventName, S3Operation};
use rustfs_utils::{
extract_params_header, extract_req_params, extract_resp_elements, get_request_host, get_request_port, get_request_user_agent,
http::headers::AMZ_REQUEST_ID,
};
use s3s::{S3Request, S3Response, S3Result};
use serde_json::Value;
use std::future::Future;
use tokio::runtime::{Builder, Handle};
use tracing::{Instrument, info_span};
/// Schedules an asynchronous task on the current runtime;
/// if there is no runtime, creates a minimal runtime execution on a new thread.
fn spawn_background<F>(fut: F)
pub(crate) fn spawn_background<F>(fut: F)
where
F: Future<Output = ()> + Send + 'static,
{
@@ -46,12 +51,30 @@ where
}
}
/// Spawn a background task with request context correlation.
/// Creates a child span with the request_id for tracing continuity,
/// ensuring audit/notify tasks can be traced back to the original request.
pub(crate) fn spawn_background_with_context<F>(request_context: Option<RequestContext>, fut: F)
where
F: Future<Output = ()> + Send + 'static,
{
match request_context {
Some(ctx) => {
let request_id = ctx.request_id.clone();
let span = info_span!("background-task", request_id = %request_id);
spawn_background(Instrument::instrument(fut, span));
}
None => spawn_background(fut),
}
}
/// A unified helper structure for building and distributing audit logs and event notifications via RAII mode at the end of an S3 operation scope.
pub struct OperationHelper {
audit_builder: Option<AuditEntryBuilder>,
api_builder: ApiDetailsBuilder,
event_builder: Option<EventArgsBuilder>,
start_time: std::time::Instant,
request_context: Option<RequestContext>,
}
impl OperationHelper {
@@ -95,18 +118,21 @@ impl OperationHelper {
api_builder = api_builder.object(&object_key);
}
// Audit builder
let mut audit_builder = AuditEntryBuilder::new("1.0", event, trigger, ApiDetails::default())
// Resolve canonical request context and request_id in a single pass:
// RequestContext.request_id > extract_request_id_from_headers() > "unknown"
let request_context = request_context_from_req(req);
let request_id = request_context
.as_ref()
.map(|ctx| ctx.request_id.clone())
.unwrap_or_else(|| extract_request_id_from_headers(&req.headers));
let audit_builder = AuditEntryBuilder::new("1.0", event, trigger, ApiDetails::default())
.remote_host(remote_host)
.user_agent(get_request_user_agent(&req.headers))
.req_host(get_request_host(&req.headers))
.req_path(req.uri.path().to_string())
.req_query(extract_req_params(req));
if let Some(req_id) = req.headers.get("x-amz-request-id")
&& let Ok(id_str) = req_id.to_str()
{
audit_builder = audit_builder.request_id(id_str);
}
.req_query(extract_req_params(req))
.request_id(&request_id);
let event_object = ObjectInfo {
bucket: bucket.clone(),
@@ -115,6 +141,12 @@ impl OperationHelper {
};
let mut req_params = extract_params_header(&req.headers);
// Inject x-amz-request-id from RequestContext into req_params for event correlation
if let Some(ref ctx) = request_context {
req_params
.entry(AMZ_REQUEST_ID.to_string())
.or_insert_with(|| ctx.x_amz_request_id.clone());
}
if let Some(principal_id) = req_info
.and_then(|info| info.cred.as_ref())
.map(|cred| cred.access_key.clone())
@@ -141,7 +173,11 @@ impl OperationHelper {
audit_builder: Some(audit_builder),
api_builder,
event_builder: Some(event_builder),
start_time: std::time::Instant::now(),
start_time: request_context
.as_ref()
.map(|ctx| ctx.start_time)
.unwrap_or_else(std::time::Instant::now),
request_context,
}
}
@@ -211,6 +247,20 @@ impl OperationHelper {
final_builder = final_builder.access_key(&sk);
}
// Inject OpenTelemetry trace context into audit tags for distributed tracing correlation
if let Some(ref ctx) = self.request_context
&& (ctx.trace_id.is_some() || ctx.span_id.is_some())
{
let mut tags = HashMap::new();
if let Some(ref tid) = ctx.trace_id {
tags.insert("traceId".to_string(), Value::String(tid.clone()));
}
if let Some(ref sid) = ctx.span_id {
tags.insert("spanId".to_string(), Value::String(sid.clone()));
}
final_builder = final_builder.tags(tags);
}
self.audit_builder = Some(final_builder);
self.api_builder = ApiDetailsBuilder(api_details); // Store final details for Drop use
}
@@ -234,7 +284,8 @@ impl Drop for OperationHelper {
fn drop(&mut self) {
// Distribute audit logs
if let Some(builder) = self.audit_builder.take() {
spawn_background(async move {
let ctx = self.request_context.clone();
spawn_background_with_context(ctx, async move {
AuditLogger::log(builder.build()).await;
});
}
@@ -246,7 +297,8 @@ impl Drop for OperationHelper {
let event_args = builder.build();
// Avoid generating notifications for copy requests
if !event_args.is_replication_request() {
spawn_background(async move {
let ctx = self.request_context.clone();
spawn_background_with_context(ctx, async move {
notifier_global::notify(event_args).await;
});
}
@@ -305,4 +357,63 @@ mod tests {
assert_eq!(event_args.version_id, "version-123");
assert_eq!(event_args.req_params.get("principalId").map(String::as_str), Some("notifyTag"));
}
#[test]
fn operation_helper_prioritizes_request_context_for_request_id() {
let input = DeleteObjectTaggingInput::builder()
.bucket("test-bucket".to_string())
.key("test-key".to_string())
.build()
.unwrap();
let mut req = build_request(input, Method::DELETE, Uri::from_static("/test-bucket/test-key"));
req.headers.insert("host", HeaderValue::from_static("example.com"));
req.headers.insert("user-agent", HeaderValue::from_static("rustfs-test"));
// Insert RequestContext (set by ingress layer) with a specific request_id
req.extensions.insert(RequestContext {
request_id: "ingress-canonical-uuid".to_string(),
x_amz_request_id: "ingress-canonical-uuid".to_string(),
trace_id: None,
span_id: None,
start_time: std::time::Instant::now(),
});
req.extensions.insert(ReqInfo {
bucket: Some("test-bucket".to_string()),
object: Some("test-key".to_string()),
..Default::default()
});
let helper = OperationHelper::new(&req, EventName::ObjectAccessedGet, S3Operation::GetObject);
// Verify the helper stored the RequestContext
assert!(helper.request_context.is_some());
assert_eq!(helper.request_context.as_ref().unwrap().request_id, "ingress-canonical-uuid");
}
#[test]
fn operation_helper_no_request_context_when_absent() {
let input = DeleteObjectTaggingInput::builder()
.bucket("test-bucket".to_string())
.key("test-key".to_string())
.build()
.unwrap();
let mut req = build_request(input, Method::DELETE, Uri::from_static("/test-bucket/test-key"));
req.headers.insert("host", HeaderValue::from_static("example.com"));
req.headers.insert("user-agent", HeaderValue::from_static("rustfs-test"));
req.headers
.insert("x-amz-request-id", HeaderValue::from_static("amz-header-uuid"));
// No RequestContext inserted
req.extensions.insert(ReqInfo {
bucket: Some("test-bucket".to_string()),
object: Some("test-key".to_string()),
..Default::default()
});
let helper = OperationHelper::new(&req, EventName::ObjectAccessedGet, S3Operation::GetObject);
// Verify the helper has no RequestContext
assert!(helper.request_context.is_none());
}
}
+1 -1
View File
@@ -21,7 +21,7 @@ pub(crate) mod entity;
pub(crate) mod helper;
pub mod lock_optimizer;
pub mod options;
pub(crate) mod readers;
pub mod request_context;
pub mod rpc;
pub(crate) mod s3_api;
mod sse;
-55
View File
@@ -1,55 +0,0 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use tokio::io::{AsyncRead, AsyncSeek};
/// Seekable in-memory async reader used by internal S3 API fast paths (e.g., GET/HEAD)
/// and by SSE flows that need a rewindable in-memory stream.
pub(crate) struct InMemoryAsyncReader {
cursor: std::io::Cursor<Vec<u8>>,
}
impl InMemoryAsyncReader {
pub(crate) fn new(data: Vec<u8>) -> Self {
Self {
cursor: std::io::Cursor::new(data),
}
}
}
impl AsyncRead for InMemoryAsyncReader {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let unfilled = buf.initialize_unfilled();
let bytes_read = std::io::Read::read(&mut self.cursor, unfilled)?;
buf.advance(bytes_read);
std::task::Poll::Ready(Ok(()))
}
}
impl AsyncSeek for InMemoryAsyncReader {
fn start_seek(mut self: std::pin::Pin<&mut Self>, position: std::io::SeekFrom) -> std::io::Result<()> {
// std::io::Cursor natively supports negative SeekCurrent offsets
// It will automatically handle validation and return an error if the final position would be negative
std::io::Seek::seek(&mut self.cursor, position)?;
Ok(())
}
fn poll_complete(self: std::pin::Pin<&mut Self>, _cx: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<u64>> {
std::task::Poll::Ready(Ok(self.cursor.position()))
}
}
+176
View File
@@ -0,0 +1,176 @@
// Copyright 2024 RustFS Team
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! Canonical request context carried through the entire request lifecycle.
//!
//! # Architecture
//!
//! ```text
//! HTTP Ingress (SetRequestIdLayer)
//! → generates x-request-id UUID
//! → RequestContextLayer creates RequestContext
//! → stores in request.extensions()
//! → sets x-amz-request-id header
//! Auth (FS::check)
//! → copies RequestContext into ReqInfo.request_context
//! Storage (FS methods)
//! → reads ReqInfo for bucket/object/version
//! → reads RequestContext for request_id/trace_id/span_id
//! Timeout Wrapper
//! → receives canonical request_id from caller
//! → passes to deadlock_detector.register_request()
//! OperationHelper
//! → reads RequestContext.request_id for audit log
//! → spawn_background_with_context() for audit/notify
//! tokio::spawn (request-internal)
//! → spawn_traced() = tokio::spawn + .instrument(Span::current())
//! ```
//!
//! # Frozen Rules (T00 Guardrails)
//!
//! ## request-id
//! - Canonical source: HTTP ingress `x-request-id` header (set by `SetRequestIdLayer`)
//! - `x-amz-request_id` is an alias for S3 compatibility, always equal to `request_id`
//! - Internal modules MUST NOT generate a second request-id under the name `request_id`
//! - Internal identifiers for sub-operations should use `operation_id` or `subtask_id`
//!
//! ## tokio::spawn usage
//! - **Request-internal tasks** (cache invalidation, metrics, read/write subtasks):
//! Use `spawn_traced()` which wraps `tokio::spawn` with `.instrument(Span::current())`
//! - **Post-request side effects** (audit flush, notify, replication enqueue):
//! Use `spawn_background_with_context()` which creates a correlated child span
//! with explicit `request_id`
//! - **Infrastructure tasks** (server loop, TLS reload, deadlock detection):
//! Plain `tokio::spawn` is acceptable; these are not request-scoped
//! - NEVER use bare `tokio::spawn` in request-handling code paths
use http::HeaderMap;
use rustfs_utils::http::headers::AMZ_REQUEST_ID;
use std::time::Instant;
/// Canonical request context carried through the entire request lifecycle.
///
/// Created exactly once at HTTP ingress. Cloned by value; never mutated after creation.
#[derive(Clone, Debug)]
pub struct RequestContext {
/// Canonical request ID (from `x-request-id` header, set by `SetRequestIdLayer`).
pub request_id: String,
/// S3-compatible request ID alias (preserves upstream `x-amz-request-id` if present,
/// otherwise equals `request_id`).
pub x_amz_request_id: String,
/// OpenTelemetry trace ID (if present from upstream propagation).
pub trace_id: Option<String>,
/// OpenTelemetry span ID (if present from upstream propagation).
pub span_id: Option<String>,
/// Request ingress timestamp.
pub start_time: Instant,
}
impl RequestContext {
/// Create a fallback `RequestContext` for paths that bypass HTTP ingress.
/// Generates a `req-{uuid}` format request-id.
pub fn fallback() -> Self {
let id = format!("req-{}", &uuid::Uuid::new_v4().to_string()[..8]);
Self {
request_id: id.clone(),
x_amz_request_id: id,
trace_id: None,
span_id: None,
start_time: Instant::now(),
}
}
}
/// Extract the canonical request ID from HTTP headers.
///
/// Priority:
/// 1. `x-request-id` (primary, set by `SetRequestIdLayer`)
/// 2. `x-amz-request-id` (fallback, from S3 client forwarding)
/// 3. `"unknown"` (no header present)
pub fn extract_request_id_from_headers(headers: &HeaderMap) -> String {
headers
.get("x-request-id")
.and_then(|v| v.to_str().ok())
.map(String::from)
.or_else(|| headers.get(AMZ_REQUEST_ID).and_then(|v| v.to_str().ok()).map(String::from))
.unwrap_or_else(|| "unknown".to_string())
}
/// Spawn a request-internal task that inherits the current tracing span.
///
/// Use this for tasks that are part of the request processing pipeline
/// (e.g., cache invalidation, metrics recording, read/write subtasks).
///
/// # Rules
/// - Do NOT use this for post-request side effects (audit, notify).
/// Use `crate::storage::helper::spawn_background_with_context` instead.
/// - Do NOT use bare `tokio::spawn` in request-handling code paths.
pub fn spawn_traced<F>(fut: F)
where
F: std::future::Future<Output = ()> + Send + 'static,
{
tokio::spawn(tracing::Instrument::instrument(fut, tracing::Span::current()));
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_request_context_clone_send_sync() {
fn assert_clone_send_sync<T: Clone + Send + Sync>() {}
assert_clone_send_sync::<RequestContext>();
}
#[test]
fn test_request_context_fallback_generates_id() {
let ctx = RequestContext::fallback();
assert!(ctx.request_id.starts_with("req-"));
assert_eq!(ctx.request_id, ctx.x_amz_request_id);
assert!(ctx.trace_id.is_none());
assert!(ctx.span_id.is_none());
}
#[test]
fn test_extract_request_id_from_x_request_id() {
let mut headers = HeaderMap::new();
headers.insert("x-request-id", "test-uuid-123".parse().unwrap());
let id = extract_request_id_from_headers(&headers);
assert_eq!(id, "test-uuid-123");
}
#[test]
fn test_extract_request_id_fallback_to_amz() {
let mut headers = HeaderMap::new();
headers.insert("x-amz-request-id", "amz-uuid-456".parse().unwrap());
let id = extract_request_id_from_headers(&headers);
assert_eq!(id, "amz-uuid-456");
}
#[test]
fn test_extract_request_id_priority() {
let mut headers = HeaderMap::new();
headers.insert("x-request-id", "x-req-789".parse().unwrap());
headers.insert("x-amz-request-id", "amz-req-000".parse().unwrap());
let id = extract_request_id_from_headers(&headers);
assert_eq!(id, "x-req-789");
}
#[test]
fn test_extract_request_id_no_headers() {
let headers = HeaderMap::new();
let id = extract_request_id_from_headers(&headers);
assert_eq!(id, "unknown");
}
}
+104
View File
@@ -14,6 +14,22 @@
use super::*;
fn lock_result_from_response(response: rustfs_lock::LockResponse) -> GenerallyLockResult {
GenerallyLockResult {
success: response.success,
error_info: response.error,
lock_info: response.lock_info.and_then(|info| serde_json::to_string(&info).ok()),
}
}
fn lock_result_from_error(error: impl Into<String>) -> GenerallyLockResult {
GenerallyLockResult {
success: false,
error_info: Some(error.into()),
lock_info: None,
}
}
impl NodeService {
pub(super) async fn handle_refresh(
&self,
@@ -144,4 +160,92 @@ impl NodeService {
})),
}
}
pub(super) async fn handle_lock_batch(
&self,
request: Request<BatchGenerallyLockRequest>,
) -> Result<Response<BatchGenerallyLockResponse>, Status> {
let request = request.into_inner();
let mut results = vec![lock_result_from_error("request was not processed"); request.args.len()];
let mut valid_requests = Vec::with_capacity(request.args.len());
let mut valid_indices = Vec::with_capacity(request.args.len());
for (idx, arg) in request.args.iter().enumerate() {
match serde_json::from_str::<LockRequest>(arg) {
Ok(args) => {
valid_requests.push(args);
valid_indices.push(idx);
}
Err(err) => {
results[idx] = lock_result_from_error(format!("can not decode args, err: {err}"));
}
}
}
if !valid_requests.is_empty() {
let lock_client = self.get_lock_client()?;
match lock_client.acquire_locks_batch(&valid_requests).await {
Ok(batch_results) => {
for (result_idx, response) in batch_results.into_iter().enumerate() {
if let Some(request_idx) = valid_indices.get(result_idx) {
results[*request_idx] = lock_result_from_response(response);
}
}
}
Err(err) => {
for request_idx in valid_indices {
results[request_idx] = lock_result_from_error(format!("can not batch lock, err: {err}"));
}
}
}
}
Ok(Response::new(BatchGenerallyLockResponse { results }))
}
pub(super) async fn handle_un_lock_batch(
&self,
request: Request<BatchGenerallyLockRequest>,
) -> Result<Response<BatchGenerallyLockResponse>, Status> {
let request = request.into_inner();
let mut results = vec![lock_result_from_error("request was not processed"); request.args.len()];
let mut lock_ids = Vec::with_capacity(request.args.len());
let mut valid_indices = Vec::with_capacity(request.args.len());
for (idx, arg) in request.args.iter().enumerate() {
match serde_json::from_str::<LockRequest>(arg) {
Ok(args) => {
lock_ids.push(args.lock_id);
valid_indices.push(idx);
}
Err(err) => {
results[idx] = lock_result_from_error(format!("can not decode args, err: {err}"));
}
}
}
if !lock_ids.is_empty() {
let lock_client = self.get_lock_client()?;
match lock_client.release_locks_batch(&lock_ids).await {
Ok(batch_results) => {
for (result_idx, success) in batch_results.into_iter().enumerate() {
if let Some(request_idx) = valid_indices.get(result_idx) {
results[*request_idx] = GenerallyLockResult {
success,
error_info: None,
lock_info: None,
};
}
}
}
Err(err) => {
for request_idx in valid_indices {
results[request_idx] = lock_result_from_error(format!("can not batch unlock, err: {err}"));
}
}
}
}
Ok(Response::new(BatchGenerallyLockResponse { results }))
}
}
+14
View File
@@ -386,6 +386,20 @@ impl Node for NodeService {
self.handle_refresh(request).await
}
async fn lock_batch(
&self,
request: Request<BatchGenerallyLockRequest>,
) -> Result<Response<BatchGenerallyLockResponse>, Status> {
self.handle_lock_batch(request).await
}
async fn un_lock_batch(
&self,
request: Request<BatchGenerallyLockRequest>,
) -> Result<Response<BatchGenerallyLockResponse>, Status> {
self.handle_un_lock_batch(request).await
}
async fn local_storage_info(
&self,
_request: Request<LocalStorageInfoRequest>,
+208 -88
View File
@@ -23,8 +23,8 @@
//!
//! ### Unified API
//! The module provides two core functions that automatically route to the correct encryption method:
//! - `apply_encryption()` - Unified encryption entry point
//! - `apply_decryption()` - Unified decryption entry point
//! - `sse_encryption()` - Unified encryption entry point
//! - `sse_decryption()` - Unified decryption entry point
//!
//! ### Managed SSE (SSE-S3 / SSE-KMS)
//! - Keys are managed by the server-side KMS service
@@ -52,8 +52,8 @@
//! part_number: None,
//! };
//!
//! if let Some(material) = apply_encryption(request).await? {
//! reader = material.wrap_reader(reader)?;
//! if let Some(material) = sse_encryption(request).await? {
//! reader = material.wrap_reader(reader);
//! metadata.extend(material.metadata);
//! }
//!
@@ -67,8 +67,10 @@
//! part_number: None,
//! };
//!
//! if let Some(material) = apply_decryption(request).await? {
//! reader = material.wrap_reader(reader)?;
//! if let Some(material) = sse_decryption(request).await? {
//! let (decrypted_reader, plaintext_size) = material.wrap_reader(reader, actual_size).await?;
//! reader = decrypted_reader;
//! content_size = plaintext_size;
//! }
//! ```
@@ -87,19 +89,17 @@ use rustfs_kms::{
service_manager::get_global_encryption_service,
types::{EncryptionMetadata, ObjectEncryptionContext},
};
use rustfs_rio::{DecryptReader, EncryptReader, HardLimitReader, Reader, WarpReader};
use rustfs_rio::{DecryptReader, DynReader, EncryptReader, HardLimitReader, ReadStream, boxed_reader, wrap_reader};
use rustfs_utils::get_env_opt_str;
use s3s::S3ErrorCode;
use s3s::dto::ServerSideEncryption;
use std::collections::HashMap;
use std::sync::{Arc, OnceLock};
use tokio::io::AsyncRead;
use tracing::{debug, error};
const INTERNAL_ENCRYPTION_KEY_ID_HEADER: &str = "x-rustfs-encryption-key-id";
use crate::error::ApiError;
use crate::storage::readers::InMemoryAsyncReader;
use rustfs_ecstore::bucket::metadata_sys;
use rustfs_ecstore::error::Error;
use s3s::dto::{SSECustomerAlgorithm, SSECustomerKey, SSECustomerKeyMD5, SSEKMSKeyId};
@@ -619,7 +619,7 @@ impl EncryptionMaterial {
/// Wrap a reader with encryption
pub fn wrap_reader<R>(&self, reader: R) -> Box<EncryptReader<R>>
where
R: Reader + 'static,
R: rustfs_rio::ReadStream + 'static,
{
Box::new(EncryptReader::new(reader, self.key_bytes, self.nonce))
}
@@ -630,42 +630,40 @@ impl DecryptionMaterial {
/// For multipart objects, use `wrap_multipart_stream` instead
pub fn wrap_single_reader<R>(&self, reader: R) -> Box<DecryptReader<R>>
where
R: Reader + 'static,
R: rustfs_rio::ReadStream + 'static,
{
Box::new(DecryptReader::new(reader, self.key_bytes, self.nonce))
}
/// Wrap a stream with multipart decryption
/// Returns the decrypted reader and the total plaintext size
pub async fn wrap_multipart_stream(
&self,
encrypted_stream: Box<dyn AsyncRead + Unpin + Send + Sync>,
) -> Result<(Box<dyn Reader>, i64), StorageError> {
pub async fn wrap_multipart_stream<R>(&self, encrypted_stream: R) -> Result<(DynReader, i64), StorageError>
where
R: ReadStream + 'static,
{
decrypt_multipart_managed_stream(encrypted_stream, &self.parts, self.key_bytes, self.nonce).await
}
/// Unified method to wrap stream with decryption and hard limit
/// Handles both single-part and multipart objects, applies decryption and size limiting
/// Accepts AsyncRead stream (from object storage) and returns (decrypted_reader, plaintext_size)
pub async fn wrap_reader(
self,
stream: Box<dyn AsyncRead + Unpin + Send + Sync>,
actual_size: i64,
) -> Result<(Box<dyn Reader>, i64), StorageError> {
let (mut final_stream, response_content_length): (Box<dyn Reader>, i64) = if self.is_multipart {
/// Accepts a readable stream (from object storage) and returns (decrypted_reader, plaintext_size)
pub async fn wrap_reader<R>(self, stream: R, actual_size: i64) -> Result<(DynReader, i64), StorageError>
where
R: ReadStream + 'static,
{
let (mut final_stream, response_content_length): (DynReader, i64) = if self.is_multipart {
// Multipart decryption
let (decrypted_reader, plain_size) = self.wrap_multipart_stream(stream).await?;
(decrypted_reader, plain_size)
} else {
// Single-part decryption - wrap AsyncRead into Reader first
let warp_reader = WarpReader::new(stream);
let decrypt_reader = self.wrap_single_reader(warp_reader);
// Single-part decryption keeps Reader capabilities via the generic wrapper helper.
let decrypt_reader = self.wrap_single_reader(wrap_reader(stream));
let plain_size = self.original_size.unwrap_or(actual_size);
(decrypt_reader, plain_size)
};
// Add hard limit reader to prevent over-reading
// final_stream is already Box<dyn Reader>, no need to wrap with WarpReader
// final_stream is already a DynReader, no need to wrap with WarpReader
let limit_reader = HardLimitReader::new(final_stream, response_content_length);
final_stream = Box::new(limit_reader);
@@ -711,8 +709,8 @@ impl DecryptionMaterial {
/// part_number: None,
/// };
///
/// if let Some(material) = apply_encryption(request).await? {
/// reader = material.wrap_reader(reader)?;
/// if let Some(material) = sse_encryption(request).await? {
/// reader = material.wrap_reader(reader);
/// metadata.extend(material.metadata);
/// }
/// ```
@@ -846,8 +844,10 @@ pub async fn sse_prepare_encryption(request: PrepareEncryptionRequest<'_>) -> Re
/// part_number: None,
/// };
///
/// if let Some(material) = apply_decryption(request).await? {
/// reader = material.wrap_reader(reader)?;
/// if let Some(material) = sse_decryption(request).await? {
/// let (decrypted_reader, plaintext_size) = material.wrap_reader(reader, actual_size).await?;
/// reader = decrypted_reader;
/// content_size = plaintext_size;
/// }
/// ```
pub async fn sse_decryption(request: DecryptionRequest<'_>) -> Result<Option<DecryptionMaterial>, ApiError> {
@@ -1642,49 +1642,43 @@ pub fn strip_managed_encryption_metadata(metadata: &mut HashMap<String, String>)
// Multipart Encryption Support
// ============================================================================
/// Derive a unique nonce for each part in a multipart upload
///
/// Uses the base nonce and increments the counter portion by part number.
/// This ensures each part has a unique nonce while maintaining determinism.
pub fn derive_part_nonce(base: [u8; 12], part_number: usize) -> [u8; 12] {
let mut nonce = base;
let current = u32::from_be_bytes([nonce[8], nonce[9], nonce[10], nonce[11]]);
let incremented = current.wrapping_add(part_number as u32);
nonce[8..12].copy_from_slice(&incremented.to_be_bytes());
nonce
derive_nonce_offset(base, 4, part_number)
}
pub(crate) async fn decrypt_multipart_managed_stream(
mut encrypted_stream: Box<dyn AsyncRead + Unpin + Send + Sync>,
#[cfg(test)]
fn derive_legacy_part_nonce(base: [u8; 12], part_number: usize) -> [u8; 12] {
derive_nonce_offset(base, 8, part_number)
}
fn derive_nonce_offset(mut base: [u8; 12], start: usize, offset: usize) -> [u8; 12] {
let current = u32::from_be_bytes([base[start], base[start + 1], base[start + 2], base[start + 3]]);
let incremented = current.wrapping_add(offset as u32);
base[start..start + 4].copy_from_slice(&incremented.to_be_bytes());
base
}
pub(crate) async fn decrypt_multipart_managed_stream<R>(
encrypted_stream: R,
parts: &[ObjectPartInfo],
key_bytes: [u8; 32],
base_nonce: [u8; 12],
) -> Result<(Box<dyn Reader>, i64), StorageError> {
let total_plain_capacity: usize = parts.iter().map(|part| part.actual_size.max(0) as usize).sum();
) -> Result<(DynReader, i64), StorageError>
where
R: ReadStream + 'static,
{
let total_plain_size = parts
.iter()
.map(|part| {
if part.actual_size > 0 {
part.actual_size
} else {
part.size as i64
}
})
.sum();
let mut plaintext = Vec::with_capacity(total_plain_capacity);
for part in parts {
if part.size == 0 {
continue;
}
let mut encrypted_part = vec![0u8; part.size];
tokio::io::AsyncReadExt::read_exact(&mut encrypted_stream, &mut encrypted_part)
.await
.map_err(|e| StorageError::other(format!("failed to read encrypted multipart segment {}: {}", part.number, e)))?;
let part_nonce = derive_part_nonce(base_nonce, part.number);
let cursor = std::io::Cursor::new(encrypted_part);
let mut decrypt_reader = DecryptReader::new(WarpReader::new(cursor), key_bytes, part_nonce);
tokio::io::AsyncReadExt::read_to_end(&mut decrypt_reader, &mut plaintext)
.await
.map_err(|e| StorageError::other(format!("failed to decrypt multipart segment {}: {}", part.number, e)))?;
}
let total_plain_size = plaintext.len() as i64;
let reader = Box::new(WarpReader::new(InMemoryAsyncReader::new(plaintext))) as Box<dyn Reader>;
let reader = boxed_reader(DecryptReader::new_multipart(wrap_reader(encrypted_stream), key_bytes, base_nonce));
Ok((reader, total_plain_size))
}
@@ -1951,13 +1945,139 @@ mod tests {
let part1 = derive_part_nonce(base, 1);
let part2 = derive_part_nonce(base, 2);
// First 8 bytes should be unchanged
assert_eq!(&base[..8], &part1[..8]);
assert_eq!(&base[..8], &part2[..8]);
assert_eq!(&base[..4], &part1[..4]);
assert_eq!(&base[8..], &part1[8..]);
assert_ne!(&base[4..8], &part1[4..8]);
assert_ne!(&part1[4..8], &part2[4..8]);
}
// Last 4 bytes should be incremented
assert_ne!(&base[8..], &part1[8..]);
assert_ne!(&part1[8..], &part2[8..]);
#[tokio::test]
async fn test_decrypt_multipart_managed_stream_accepts_legacy_part_nonce_layout() {
use std::io::Cursor;
use tokio::io::AsyncReadExt;
let key_bytes = [7u8; 32];
let base_nonce = [3u8; 12];
let part_one_plaintext = vec![0x11; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE + 19];
let part_two_plaintext = vec![0x22; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE + 37];
let part_one_nonce = derive_legacy_part_nonce(base_nonce, 1);
let part_two_nonce = derive_legacy_part_nonce(base_nonce, 2);
let first_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_one_plaintext.clone()), key_bytes, part_one_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let second_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_two_plaintext.clone()), key_bytes, part_two_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let mut encrypted_stream = Vec::with_capacity(first_part.len() + second_part.len());
encrypted_stream.extend_from_slice(&first_part);
encrypted_stream.extend_from_slice(&second_part);
let parts = vec![
ObjectPartInfo {
number: 1,
size: first_part.len(),
actual_size: part_one_plaintext.len() as i64,
..Default::default()
},
ObjectPartInfo {
number: 2,
size: second_part.len(),
actual_size: part_two_plaintext.len() as i64,
..Default::default()
},
];
let (mut decrypted_reader, plaintext_size) =
decrypt_multipart_managed_stream(Cursor::new(encrypted_stream), &parts, key_bytes, base_nonce)
.await
.unwrap();
let mut decrypted = Vec::new();
decrypted_reader.read_to_end(&mut decrypted).await.unwrap();
let mut expected = part_one_plaintext;
expected.extend_from_slice(&part_two_plaintext);
assert_eq!(plaintext_size, expected.len() as i64);
assert_eq!(decrypted, expected);
}
#[tokio::test]
async fn test_decrypt_multipart_managed_stream_supports_current_nonce_layout() {
use std::io::Cursor;
use tokio::io::AsyncReadExt;
let key_bytes = [9u8; 32];
let base_nonce = [5u8; 12];
let part_one_plaintext = vec![0x33; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE + 11];
let part_two_plaintext = vec![0x44; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE * 2 + 7];
let part_one_nonce = derive_part_nonce(base_nonce, 1);
let part_two_nonce = derive_part_nonce(base_nonce, 2);
let first_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_one_plaintext.clone()), key_bytes, part_one_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let second_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_two_plaintext.clone()), key_bytes, part_two_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let mut encrypted_stream = Vec::with_capacity(first_part.len() + second_part.len());
encrypted_stream.extend_from_slice(&first_part);
encrypted_stream.extend_from_slice(&second_part);
let parts = vec![
ObjectPartInfo {
number: 1,
size: first_part.len(),
actual_size: part_one_plaintext.len() as i64,
..Default::default()
},
ObjectPartInfo {
number: 2,
size: second_part.len(),
actual_size: part_two_plaintext.len() as i64,
..Default::default()
},
];
let (mut decrypted_reader, plaintext_size) =
decrypt_multipart_managed_stream(Cursor::new(encrypted_stream), &parts, key_bytes, base_nonce)
.await
.unwrap();
let mut decrypted = Vec::new();
decrypted_reader.read_to_end(&mut decrypted).await.unwrap();
let mut expected = part_one_plaintext;
expected.extend_from_slice(&part_two_plaintext);
assert_eq!(plaintext_size, expected.len() as i64);
assert_eq!(decrypted, expected);
}
#[test]
@@ -2436,8 +2556,8 @@ mod tests {
println!("Original plaintext: {:?}", String::from_utf8_lossy(plaintext));
println!("Plaintext length: {} bytes", plaintext.len());
// 4. Encrypt with EncryptReader (wrap Cursor with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// 4. Encrypt with EncryptReader.
let plaintext_reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
// Read encrypted data
@@ -2460,8 +2580,8 @@ mod tests {
"Encrypted data should be different from plaintext"
);
// 5. Decrypt with DecryptReader (wrap Cursor with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// 5. Decrypt with DecryptReader.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
// Read decrypted data
@@ -2502,8 +2622,8 @@ mod tests {
let plaintext: Vec<u8> = (0..plaintext_size).map(|i| (i % 256) as u8).collect();
println!("Testing with {} bytes of data", plaintext.len());
// Encrypt (wrap with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.clone()));
// Encrypt.
let plaintext_reader = Cursor::new(plaintext.clone());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
let mut encrypted_data = Vec::new();
@@ -2514,8 +2634,8 @@ mod tests {
println!("Encrypted {} bytes to {} bytes", plaintext.len(), encrypted_data.len());
// Decrypt (wrap with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
let mut decrypted_data = Vec::new();
@@ -2560,14 +2680,14 @@ mod tests {
// Same plaintext
let plaintext = b"Same plaintext";
// Encrypt with first key (wrap with WarpReader)
let reader1 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with first key.
let reader1 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader1 = EncryptReader::new(reader1, data_key1.plaintext_key, data_key1.nonce);
let mut encrypted1 = Vec::new();
encrypt_reader1.read_to_end(&mut encrypted1).await.unwrap();
// Encrypt with second key (wrap with WarpReader)
let reader2 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with second key.
let reader2 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader2 = EncryptReader::new(reader2, data_key2.plaintext_key, data_key2.nonce);
let mut encrypted2 = Vec::new();
encrypt_reader2.read_to_end(&mut encrypted2).await.unwrap();
@@ -2620,14 +2740,14 @@ mod tests {
// 5. Use decrypted key to encrypt/decrypt data
let plaintext = b"Test data with decrypted DEK";
// Encrypt with original key (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with original key.
let reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(reader, original_plaintext_key, original_nonce);
let mut encrypted_data = Vec::new();
encrypt_reader.read_to_end(&mut encrypted_data).await.unwrap();
// Decrypt with recovered key (simulating GET operation) (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt with recovered key (simulating GET operation).
let reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(
reader,
decrypted_plaintext_key,
+17 -17
View File
@@ -16,7 +16,7 @@
mod tests {
use crate::storage::sse::SseDekProvider;
use crate::storage::sse::TestSseDekProvider;
use rustfs_rio::{DecryptReader, EncryptReader, WarpReader};
use rustfs_rio::{DecryptReader, EncryptReader};
use std::io::Cursor;
use tokio::io::AsyncReadExt;
@@ -51,8 +51,8 @@ mod tests {
println!("Original plaintext: {:?}", String::from_utf8_lossy(plaintext));
println!("Plaintext length: {} bytes", plaintext.len());
// Step 4: Encrypt using EncryptReader (wrap Cursor with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Step 4: Encrypt using EncryptReader.
let plaintext_reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
// Read encrypted data
@@ -75,8 +75,8 @@ mod tests {
"Encrypted data should be different from plaintext"
);
// Step 5: Decrypt using DecryptReader (wrap Cursor with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// Step 5: Decrypt using DecryptReader.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
// Read decrypted data
@@ -115,8 +115,8 @@ mod tests {
let plaintext: Vec<u8> = (0..plaintext_size).map(|i| (i % 256) as u8).collect();
println!("Testing with {} bytes of data", plaintext.len());
// Encrypt (wrap with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.clone()));
// Encrypt.
let plaintext_reader = Cursor::new(plaintext.clone());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
let mut encrypted_data = Vec::new();
@@ -127,8 +127,8 @@ mod tests {
println!("Encrypted {} bytes to {} bytes", plaintext.len(), encrypted_data.len());
// Decrypt (wrap with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
let mut decrypted_data = Vec::new();
@@ -171,14 +171,14 @@ mod tests {
// Same plaintext
let plaintext = b"Same plaintext";
// Encrypt with first key (wrap with WarpReader)
let reader1 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with first key.
let reader1 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader1 = EncryptReader::new(reader1, data_key1.plaintext_key, data_key1.nonce);
let mut encrypted1 = Vec::new();
encrypt_reader1.read_to_end(&mut encrypted1).await.unwrap();
// Encrypt with second key (wrap with WarpReader)
let reader2 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with second key.
let reader2 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader2 = EncryptReader::new(reader2, data_key2.plaintext_key, data_key2.nonce);
let mut encrypted2 = Vec::new();
encrypt_reader2.read_to_end(&mut encrypted2).await.unwrap();
@@ -226,14 +226,14 @@ mod tests {
// Step 4: Use decrypted key to encrypt/decrypt data
let plaintext = b"Test data with decrypted DEK";
// Encrypt with original key (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with original key.
let reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(reader, original_plaintext_key, original_nonce);
let mut encrypted_data = Vec::new();
encrypt_reader.read_to_end(&mut encrypted_data).await.unwrap();
// Decrypt with recovered key (simulating GET operation) (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt with recovered key (simulating GET operation).
let reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(
reader,
decrypted_plaintext_key,
+10 -7
View File
@@ -234,12 +234,15 @@ pub struct RequestTimeoutWrapper {
impl RequestTimeoutWrapper {
/// Create a new timeout wrapper with the given configuration.
///
/// Note: This uses a sentinel request_id. Prefer `with_request_id()` to pass
/// the canonical request-id from `RequestContext`.
pub fn new(config: TimeoutConfig) -> Self {
Self {
config,
start_time: Instant::now(),
cancel_token: CancellationToken::new(),
request_id: format!("req-{}", &uuid::Uuid::new_v4().to_string()[..8]),
request_id: "no-request-id".to_string(),
}
}
@@ -253,17 +256,17 @@ impl RequestTimeoutWrapper {
}
}
/// Create a new timeout wrapper with operation size for dynamic timeout calculation
/// Create a new timeout wrapper with operation size for dynamic timeout calculation.
///
/// Note: This uses a sentinel request_id. Prefer `with_request_id()` to pass
/// the canonical request-id from `RequestContext`.
pub fn with_operation_size(config: TimeoutConfig, operation_size: Option<u64>) -> Self {
// Store operation size in config for later use
// Note: Currently we don't store the size in the wrapper itself,
// but the config can be used to calculate appropriate timeout
let _ = operation_size; // Suppress unused warning for now
let _ = operation_size;
Self {
config,
start_time: Instant::now(),
cancel_token: CancellationToken::new(),
request_id: format!("req-{}", &uuid::Uuid::new_v4().to_string()[..8]),
request_id: "no-request-id".to_string(),
}
}