fix(admin): harden STS and KMS authorization checks (#2653)

Co-authored-by: houseme <housemecn@gmail.com>
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: loverustfs <hello@rustfs.com>
This commit is contained in:
Henry Guo
2026-04-25 09:34:18 +08:00
committed by GitHub
parent 94f64acc87
commit c717195de2
4 changed files with 202 additions and 51 deletions
+16 -8
View File
@@ -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<Self, Self::Error> {
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
+6 -6
View File
@@ -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::<Option<RemoteAddr>>().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::<Option<RemoteAddr>>().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::<Option<RemoteAddr>>().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::<Option<RemoteAddr>>().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::<Option<RemoteAddr>>().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::<Option<RemoteAddr>>().and_then(|opt| opt.map(|a| a.0)),
)
.await?;
+35 -8
View File
@@ -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::<Option<RemoteAddr>>().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<s3s::auth::Credentials>,
uri: http::Uri,
headers: http::HeaderMap,
remote_addr: Option<std::net::SocketAddr>,
body: AssumeRoleRequest,
) -> S3Result<S3Response<(StatusCode, Body)>> {
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: {:?}",
+145 -29
View File
@@ -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<T> = Pin<Box<dyn Stream<Item = Result<T, Status>> + 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<String> {
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<ReadAtResponse>;
async fn read_at(&self, _request: Request<Streaming<ReadAtRequest>>) -> Result<Response<Self::ReadAtStream>, Status> {
info!("read_at");
unimplemented!("read_at");
Err(unimplemented_rpc("read_at"))
}
async fn list_dir(&self, request: Request<ListDirRequest>) -> Result<Response<ListDirResponse>, Status> {
@@ -462,35 +466,35 @@ impl Node for NodeService {
&self,
_request: Request<StartProfilingRequest>,
) -> Result<Response<StartProfilingResponse>, Status> {
todo!()
Err(unimplemented_rpc("start_profiling"))
}
async fn download_profile_data(
&self,
_request: Request<DownloadProfileDataRequest>,
) -> Result<Response<DownloadProfileDataResponse>, Status> {
todo!()
Err(unimplemented_rpc("download_profile_data"))
}
async fn get_bucket_stats(
&self,
_request: Request<GetBucketStatsDataRequest>,
) -> Result<Response<GetBucketStatsDataResponse>, Status> {
todo!()
Err(unimplemented_rpc("get_bucket_stats"))
}
async fn get_sr_metrics(
&self,
_request: Request<GetSrMetricsDataRequest>,
) -> Result<Response<GetSrMetricsDataResponse>, Status> {
todo!()
Err(unimplemented_rpc("get_sr_metrics"))
}
async fn get_all_bucket_stats(
&self,
_request: Request<GetAllBucketStatsRequest>,
) -> Result<Response<GetAllBucketStatsResponse>, 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<SignalServiceRequest>) -> Result<Response<SignalServiceResponse>, 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<BackgroundHealStatusRequest>,
) -> Result<Response<BackgroundHealStatusResponse>, Status> {
todo!()
Err(unimplemented_rpc("background_heal_status"))
}
async fn get_metacache_listing(
&self,
_request: Request<GetMetacacheListingRequest>,
) -> Result<Response<GetMetacacheListingResponse>, Status> {
todo!()
Err(unimplemented_rpc("get_metacache_listing"))
}
async fn update_metacache_listing(
&self,
_request: Request<UpdateMetacacheListingRequest>,
) -> Result<Response<UpdateMetacacheListingResponse>, Status> {
todo!()
Err(unimplemented_rpc("update_metacache_listing"))
}
async fn reload_pool_meta(
@@ -894,7 +894,7 @@ impl Node for NodeService {
&self,
_request: Request<LoadTransitionTierConfigRequest>,
) -> Result<Response<LoadTransitionTierConfigResponse>, 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<T>(response: Result<Response<T>, 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<tonic::transport::Channel> {
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() {