From c717195de2d03ffd1ffd25690a4c4b8d3f9d1606 Mon Sep 17 00:00:00 2001 From: Henry Guo Date: Sat, 25 Apr 2026 09:34:18 +0800 Subject: [PATCH] fix(admin): harden STS and KMS authorization checks (#2653) Co-authored-by: houseme Co-authored-by: Claude Sonnet 4.6 Co-authored-by: loverustfs --- crates/policy/src/policy/action.rs | 24 ++-- rustfs/src/admin/handlers/kms_keys.rs | 12 +- rustfs/src/admin/handlers/sts.rs | 43 ++++-- rustfs/src/storage/rpc/node_service.rs | 174 ++++++++++++++++++++----- 4 files changed, 202 insertions(+), 51 deletions(-) diff --git a/crates/policy/src/policy/action.rs b/crates/policy/src/policy/action.rs index 45c39a452..342cba397 100644 --- a/crates/policy/src/policy/action.rs +++ b/crates/policy/src/policy/action.rs @@ -599,15 +599,13 @@ impl AdminAction { } } -#[derive(Serialize, Deserialize, Hash, PartialEq, Eq, Clone, IntoStaticStr, Debug, Copy)] +#[derive(Serialize, Deserialize, Hash, PartialEq, Eq, Clone, IntoStaticStr, Debug, Copy, EnumString)] #[serde(try_from = "&str", into = "&str")] -pub enum StsAction {} - -impl TryFrom<&str> for StsAction { - type Error = strum::ParseError; - fn try_from(_value: &str) -> std::result::Result { - Err(strum::ParseError::VariantNotFound) - } +pub enum StsAction { + #[strum(serialize = "sts:*")] + AllActions, + #[strum(serialize = "sts:AssumeRole")] + AssumeRoleAction, } #[derive(Serialize, Deserialize, Hash, PartialEq, Eq, Clone, IntoStaticStr, Debug, Copy, EnumString)] @@ -629,6 +627,16 @@ mod tests { assert!(matches!(action, Action::S3Action(S3Action::AllActions))); } + #[test] + fn test_sts_action_parsing() { + let action = Action::try_from("sts:AssumeRole").expect("Should parse STS AssumeRole action"); + assert!(matches!(action, Action::StsAction(StsAction::AssumeRoleAction))); + + let wildcard = Action::try_from("sts:*").expect("Should parse STS wildcard action"); + assert!(matches!(wildcard, Action::StsAction(StsAction::AllActions))); + assert!(wildcard.is_match(&action)); + } + #[test] fn test_actionset_serialize_single_element() { // Single element should serialize as array for S3 specification compliance diff --git a/rustfs/src/admin/handlers/kms_keys.rs b/rustfs/src/admin/handlers/kms_keys.rs index 25bb7b1dc..98a6b5d44 100644 --- a/rustfs/src/admin/handlers/kms_keys.rs +++ b/rustfs/src/admin/handlers/kms_keys.rs @@ -170,7 +170,7 @@ impl Operation for CreateKeyHandler { &cred, owner, false, - vec![Action::AdminAction(AdminAction::ServerInfoAdminAction)], // TODO: Add specific KMS action + vec![Action::AdminAction(AdminAction::KMSCreateKeyAdminAction)], req.extensions.get::>().and_then(|opt| opt.map(|a| a.0)), ) .await?; @@ -249,7 +249,7 @@ impl Operation for DescribeKeyHandler { &cred, owner, false, - vec![Action::AdminAction(AdminAction::ServerInfoAdminAction)], + vec![Action::AdminAction(AdminAction::KMSKeyStatusAdminAction)], req.extensions.get::>().and_then(|opt| opt.map(|a| a.0)), ) .await?; @@ -351,7 +351,7 @@ impl Operation for ListKeysHandler { &cred, owner, false, - vec![Action::AdminAction(AdminAction::ServerInfoAdminAction)], + vec![Action::AdminAction(AdminAction::KMSKeyStatusAdminAction)], req.extensions.get::>().and_then(|opt| opt.map(|a| a.0)), ) .await?; @@ -479,7 +479,7 @@ impl Operation for CreateKmsKeyHandler { &cred, owner, false, - vec![Action::AdminAction(AdminAction::ServerInfoAdminAction)], + vec![Action::AdminAction(AdminAction::KMSCreateKeyAdminAction)], req.extensions.get::>().and_then(|opt| opt.map(|a| a.0)), ) .await?; @@ -891,7 +891,7 @@ impl Operation for ListKmsKeysHandler { &cred, owner, false, - vec![Action::AdminAction(AdminAction::ServerInfoAdminAction)], + vec![Action::AdminAction(AdminAction::KMSKeyStatusAdminAction)], req.extensions.get::>().and_then(|opt| opt.map(|a| a.0)), ) .await?; @@ -1003,7 +1003,7 @@ impl Operation for DescribeKmsKeyHandler { &cred, owner, false, - vec![Action::AdminAction(AdminAction::ServerInfoAdminAction)], + vec![Action::AdminAction(AdminAction::KMSKeyStatusAdminAction)], req.extensions.get::>().and_then(|opt| opt.map(|a| a.0)), ) .await?; diff --git a/rustfs/src/admin/handlers/sts.rs b/rustfs/src/admin/handlers/sts.rs index 2cf9bd522..bc72f2a92 100644 --- a/rustfs/src/admin/handlers/sts.rs +++ b/rustfs/src/admin/handlers/sts.rs @@ -20,6 +20,7 @@ use crate::{ }, auth::{check_key_valid, extract_string_list_claim, get_session_token}, server::ADMIN_PREFIX, + server::RemoteAddr, }; use http::StatusCode; use http::header::HeaderValue; @@ -30,7 +31,13 @@ use rustfs_credentials::get_global_action_cred; use rustfs_ecstore::bucket::utils::serialize; use rustfs_iam::{manager::get_token_signing_key, oidc::OidcClaims, sys::SESSION_POLICY_NAME}; use rustfs_madmin::{SITE_REPL_API_VERSION, SRIAMItem, SRSTSCredential}; -use rustfs_policy::{auth::get_new_credentials_with_metadata, policy::Policy}; +use rustfs_policy::{ + auth::get_new_credentials_with_metadata, + policy::{ + Args, Policy, + action::{Action, StsAction}, + }, +}; use s3s::{ Body, S3Error, S3ErrorCode, S3Request, S3Response, S3Result, dto::{AssumeRoleOutput, Credentials, Timestamp}, @@ -147,7 +154,10 @@ impl Operation for AssumeRoleHandle { let body: AssumeRoleRequest = from_bytes(&bytes).map_err(|_e| s3_error!(InvalidRequest, "invalid STS request format"))?; match body.action.as_str() { - ASSUME_ROLE_ACTION => handle_assume_role(req.credentials, req.uri, req.headers, body).await, + ASSUME_ROLE_ACTION => { + let remote_addr = req.extensions.get::>().and_then(|opt| opt.map(|a| a.0)); + handle_assume_role(req.credentials, req.uri, req.headers, remote_addr, body).await + } ASSUME_ROLE_WITH_WEB_IDENTITY_ACTION => handle_assume_role_with_web_identity(body).await, _ => Err(s3_error!(InvalidArgument, "unsupported Action")), } @@ -159,6 +169,7 @@ async fn handle_assume_role( credentials: Option, uri: http::Uri, headers: http::HeaderMap, + remote_addr: Option, body: AssumeRoleRequest, ) -> S3Result> { let Some(user) = credentials else { @@ -170,13 +181,33 @@ async fn handle_assume_role( return Err(s3_error!(InvalidRequest, "AccessDenied1")); } - let (cred, _owner) = check_key_valid(get_session_token(&uri, &headers).unwrap_or_default(), &user.access_key).await?; + let (cred, owner) = check_key_valid(get_session_token(&uri, &headers).unwrap_or_default(), &user.access_key).await?; - // TODO: Check permissions, do not allow STS access if cred.is_temp() || cred.is_service_account() { return Err(s3_error!(InvalidRequest, "AccessDenied")); } + let Ok(iam_store) = rustfs_iam::get() else { + return Err(s3_error!(InvalidRequest, "iam not init")); + }; + let conditions = crate::auth::get_condition_values(&headers, &cred, None, None, remote_addr); + if !iam_store + .is_allowed(&Args { + account: &cred.access_key, + groups: &cred.groups, + action: Action::StsAction(StsAction::AssumeRoleAction), + conditions: &conditions, + is_owner: owner, + claims: cred.claims_or_empty(), + deny_only: false, + bucket: "", + object: "", + }) + .await + { + return Err(s3_error!(AccessDenied, "Access Denied")); + } + if body.version.as_str() != ASSUME_ROLE_VERSION { return Err(s3_error!(InvalidArgument, "not support version")); } @@ -200,10 +231,6 @@ async fn handle_assume_role( claims.insert("parent".to_string(), Value::String(cred.access_key.clone())); - let Ok(iam_store) = rustfs_iam::get() else { - return Err(s3_error!(InvalidRequest, "iam not init")); - }; - if let Err(_err) = iam_store.policy_db_get(&cred.access_key, &cred.groups).await { error!( "AssumeRole get policy failed, err: {:?}, access_key: {:?}, groups: {:?}", diff --git a/rustfs/src/storage/rpc/node_service.rs b/rustfs/src/storage/rpc/node_service.rs index 53960d8ef..c322ffb4b 100644 --- a/rustfs/src/storage/rpc/node_service.rs +++ b/rustfs/src/storage/rpc/node_service.rs @@ -44,7 +44,7 @@ use rustfs_protos::{ proto_gen::node_service::{node_service_server::NodeService as Node, *}, }; use serde::{Deserialize, Serialize}; -use std::{collections::HashMap, io::Cursor, pin::Pin, sync::Arc}; +use std::{io::Cursor, pin::Pin, sync::Arc}; use tokio::spawn; use tokio::sync::mpsc; use tokio_stream::wrappers::ReceiverStream; @@ -53,6 +53,10 @@ use tracing::{debug, error, info, warn}; type ResponseStream = Pin> + Send>>; +fn unimplemented_rpc(method: &str) -> Status { + Status::unimplemented(format!("{method} is not implemented")) +} + fn background_rebalance_start_error_message(result: rustfs_ecstore::error::Result<()>) -> Option { result.err().map(|err| format!("start_rebalance failed: {err}")) } @@ -184,13 +188,13 @@ impl Node for NodeService { info!("write_stream"); let _ = request; - unimplemented!("write_stream"); + Err(unimplemented_rpc("write_stream")) } type ReadAtStream = ResponseStream; async fn read_at(&self, _request: Request>) -> Result, Status> { info!("read_at"); - unimplemented!("read_at"); + Err(unimplemented_rpc("read_at")) } async fn list_dir(&self, request: Request) -> Result, Status> { @@ -462,35 +466,35 @@ impl Node for NodeService { &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("start_profiling")) } async fn download_profile_data( &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("download_profile_data")) } async fn get_bucket_stats( &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("get_bucket_stats")) } async fn get_sr_metrics( &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("get_sr_metrics")) } async fn get_all_bucket_stats( &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("get_all_bucket_stats")) } async fn load_bucket_metadata( @@ -785,33 +789,29 @@ impl Node for NodeService { } async fn signal_service(&self, request: Request) -> Result, Status> { - let request = request.into_inner(); - let _vars = match request.vars { - Some(vars) => vars.value, - None => HashMap::new(), - }; - todo!() + let _request = request.into_inner(); + Err(unimplemented_rpc("signal_service")) } async fn background_heal_status( &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("background_heal_status")) } async fn get_metacache_listing( &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("get_metacache_listing")) } async fn update_metacache_listing( &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("update_metacache_listing")) } async fn reload_pool_meta( @@ -894,7 +894,7 @@ impl Node for NodeService { &self, _request: Request, ) -> Result, Status> { - todo!() + Err(unimplemented_rpc("load_transition_tier_config")) } } @@ -904,18 +904,23 @@ mod tests { use super::*; use Request; use rustfs_protos::proto_gen::node_service::{ - CheckPartsRequest, DeleteBucketMetadataRequest, DeleteBucketRequest, DeletePathsRequest, DeletePolicyRequest, - DeleteRequest, DeleteServiceAccountRequest, DeleteUserRequest, DeleteVersionRequest, DeleteVersionsRequest, - DeleteVolumeRequest, DiskInfoRequest, GenerallyLockRequest, GetBucketInfoRequest, GetCpusRequest, GetMemInfoRequest, - GetNetInfoRequest, GetOsInfoRequest, GetPartitionsRequest, GetProcInfoRequest, GetSeLinuxInfoRequest, - GetSysConfigRequest, GetSysErrorsRequest, HealBucketRequest, ListBucketRequest, ListDirRequest, ListVolumesRequest, - LoadBucketMetadataRequest, LoadGroupRequest, LoadPolicyMappingRequest, LoadPolicyRequest, LoadRebalanceMetaRequest, - LoadServiceAccountRequest, LoadUserRequest, LocalStorageInfoRequest, MakeBucketRequest, MakeVolumeRequest, - MakeVolumesRequest, PingRequest, ReadAllRequest, ReadMultipleRequest, ReadVersionRequest, ReadXlRequest, + BackgroundHealStatusRequest, CheckPartsRequest, DeleteBucketMetadataRequest, DeleteBucketRequest, DeletePathsRequest, + DeletePolicyRequest, DeleteRequest, DeleteServiceAccountRequest, DeleteUserRequest, DeleteVersionRequest, + DeleteVersionsRequest, DeleteVolumeRequest, DiskInfoRequest, DownloadProfileDataRequest, GenerallyLockRequest, + GetAllBucketStatsRequest, GetBucketInfoRequest, GetBucketStatsDataRequest, GetCpusRequest, GetMemInfoRequest, + GetMetacacheListingRequest, GetNetInfoRequest, GetOsInfoRequest, GetPartitionsRequest, GetProcInfoRequest, + GetSeLinuxInfoRequest, GetSrMetricsDataRequest, GetSysConfigRequest, GetSysErrorsRequest, HealBucketRequest, + ListBucketRequest, ListDirRequest, ListVolumesRequest, LoadBucketMetadataRequest, LoadGroupRequest, + LoadPolicyMappingRequest, LoadPolicyRequest, LoadRebalanceMetaRequest, LoadServiceAccountRequest, + LoadTransitionTierConfigRequest, LoadUserRequest, LocalStorageInfoRequest, MakeBucketRequest, MakeVolumeRequest, + MakeVolumesRequest, PingRequest, ReadAllRequest, ReadAtRequest, ReadMultipleRequest, ReadVersionRequest, ReadXlRequest, ReloadPoolMetaRequest, ReloadSiteReplicationConfigRequest, RenameDataRequest, RenameFileRequest, RenamePartRequest, - ServerInfoRequest, StatVolumeRequest, StopRebalanceRequest, UpdateMetadataRequest, VerifyFileRequest, WriteAllRequest, - WriteMetadataRequest, + ServerInfoRequest, SignalServiceRequest, StartProfilingRequest, StatVolumeRequest, StopRebalanceRequest, + UpdateMetacacheListingRequest, UpdateMetadataRequest, VerifyFileRequest, WriteAllRequest, WriteMetadataRequest, + WriteRequest, node_service_client::NodeServiceClient, node_service_server::NodeServiceServer, }; + use tokio::net::TcpListener; + use tokio_stream::wrappers::TcpListenerStream; fn create_test_node_service() -> NodeService { make_server() @@ -2235,7 +2240,118 @@ mod tests { assert!(reload_response.error_info.is_some()); } - // Note: signal_service test is skipped because it contains todo!() and would panic + fn assert_unimplemented_status(response: Result, Status>, method: &str) { + let err = match response { + Ok(_) => panic!("unimplemented RPC should return an error status"), + Err(err) => err, + }; + assert_eq!(err.code(), tonic::Code::Unimplemented); + assert!( + err.message().contains(method), + "expected method name in status message, got {:?}", + err.message() + ); + } + + #[tokio::test] + async fn test_unimplemented_rpcs_return_status() { + let service = create_test_node_service(); + + assert_unimplemented_status( + service.start_profiling(Request::new(StartProfilingRequest::default())).await, + "start_profiling", + ); + assert_unimplemented_status( + service + .download_profile_data(Request::new(DownloadProfileDataRequest::default())) + .await, + "download_profile_data", + ); + assert_unimplemented_status( + service + .get_bucket_stats(Request::new(GetBucketStatsDataRequest::default())) + .await, + "get_bucket_stats", + ); + assert_unimplemented_status( + service.get_sr_metrics(Request::new(GetSrMetricsDataRequest::default())).await, + "get_sr_metrics", + ); + assert_unimplemented_status( + service + .get_all_bucket_stats(Request::new(GetAllBucketStatsRequest::default())) + .await, + "get_all_bucket_stats", + ); + assert_unimplemented_status( + service.signal_service(Request::new(SignalServiceRequest::default())).await, + "signal_service", + ); + assert_unimplemented_status( + service + .background_heal_status(Request::new(BackgroundHealStatusRequest::default())) + .await, + "background_heal_status", + ); + assert_unimplemented_status( + service + .get_metacache_listing(Request::new(GetMetacacheListingRequest::default())) + .await, + "get_metacache_listing", + ); + assert_unimplemented_status( + service + .update_metacache_listing(Request::new(UpdateMetacacheListingRequest::default())) + .await, + "update_metacache_listing", + ); + assert_unimplemented_status( + service + .load_transition_tier_config(Request::new(LoadTransitionTierConfigRequest::default())) + .await, + "load_transition_tier_config", + ); + } + + async fn connect_test_node_service_client() -> NodeServiceClient { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let service = create_test_node_service(); + + tokio::spawn(async move { + tonic::transport::Server::builder() + .add_service(NodeServiceServer::new(service)) + .serve_with_incoming(TcpListenerStream::new(listener)) + .await + .unwrap(); + }); + + NodeServiceClient::connect(format!("http://{addr}")).await.unwrap() + } + + #[tokio::test] + async fn test_write_stream_unimplemented() { + let mut client = connect_test_node_service_client().await; + let request = tokio_stream::iter([WriteRequest::default()]); + + let response = client.write_stream(request).await; + + let err = response.expect_err("write_stream should return unimplemented status"); + assert_eq!(err.code(), tonic::Code::Unimplemented); + assert!(err.message().contains("write_stream")); + } + + #[tokio::test] + async fn test_read_at_unimplemented() { + let mut client = connect_test_node_service_client().await; + let request = tokio_stream::iter([ReadAtRequest::default()]); + + let response = client.read_at(request).await; + + let err = response.expect_err("read_at should return unimplemented status"); + assert_eq!(err.code(), tonic::Code::Unimplemented); + assert!(err.message().contains("read_at")); + } #[tokio::test] async fn test_node_service_debug() {