fix(aws): Assume role for Gov Cloud (#4254)

Co-authored-by: Sergio Garcia <38561120+sergargar@users.noreply.github.com>
This commit is contained in:
Pepe Fagoaga
2024-06-18 09:37:23 -04:00
committed by GitHub
co-authored by Sergio Garcia
parent 625be45742
commit e8a94733bf
5 changed files with 112 additions and 61 deletions
+14 -16
View File
@@ -28,6 +28,7 @@ from prowler.lib.mutelist.mutelist import (
)
from prowler.lib.utils.utils import open_file, parse_json_file, print_boxes
from prowler.providers.aws.config import (
AWS_REGION_US_EAST_1,
AWS_STS_GLOBAL_ENDPOINT_REGION,
BOTO3_USER_AGENT_EXTRA,
ROLE_SESSION_NAME,
@@ -144,6 +145,7 @@ class AwsProvider(Provider):
input_mfa,
input_session_duration,
input_role_session_name,
sts_region,
)
# Assume the IAM Role
logger.info(f"Assuming role: {assumed_role_information.role_arn.arn}")
@@ -188,6 +190,7 @@ class AwsProvider(Provider):
input_mfa,
input_session_duration,
input_role_session_name,
sts_region,
)
# Assume the Organizations IAM Role
logger.info(
@@ -214,11 +217,6 @@ class AwsProvider(Provider):
"Generated new session for to get the AWS Organizations metadata"
)
# TODO: Do we need to modify the identity here? I think not since it is not used
# self._identity.account = assumed_role.info.role_arn.account_id
# self._identity.partition = assumed_role.info.role_arn.partition
# self._identity.account_arn = f"arn:{self._identity.partition}:iam::{assumed_role.info.role_arn.account_id}:root"
self._organizations_metadata = self.get_organizations_info(
aws_organizations_session, self._identity.account
)
@@ -404,10 +402,9 @@ class AwsProvider(Provider):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
# TODO: This can be moved to another class since it doesn't need self
def get_profile_region(self, session: Session):
# TODO: read "us-east-1" from another place
profile_region = "us-east-1"
@staticmethod
def get_profile_region(session: Session):
profile_region = AWS_REGION_US_EAST_1
if session.region_name:
profile_region = session.region_name
@@ -424,7 +421,6 @@ class AwsProvider(Provider):
logger.info(f"Original AWS Caller Identity ARN: {caller_identity.arn}")
partition = parse_iam_credentials_arn(caller_identity.arn.arn).partition
return AWSIdentityInfo(
account=caller_identity.account,
account_arn=f"arn:{partition}:iam::{caller_identity.account}:root",
@@ -437,7 +433,10 @@ class AwsProvider(Provider):
)
def setup_session(
self, input_mfa: bool, input_profile: str, input_role: str = None
self,
input_mfa: bool,
input_profile: str,
input_role: str = None,
) -> Session:
try:
logger.info("Creating original session ...")
@@ -479,6 +478,7 @@ class AwsProvider(Provider):
input_mfa: str,
session_duration: int,
role_session_name: str,
sts_region: str = AWS_STS_GLOBAL_ENDPOINT_REGION,
) -> AWSAssumeRoleInfo:
"""
set_assumed_role_info returns a AWSAssumeRoleInfo object
@@ -490,6 +490,7 @@ class AwsProvider(Provider):
external_id=input_external_id,
mfa_enabled=input_mfa,
role_session_name=role_session_name,
sts_region=sts_region,
)
def setup_assumed_session(
@@ -823,8 +824,6 @@ class AwsProvider(Provider):
self,
session: Session,
assumed_role_info: AWSAssumeRoleInfo,
# TODO: remove I think
# sts_endpoint_region: str = None,
) -> AWSCredentials:
"""
assume_role assumes the IAM roles passed with the given session and returns AWSCredentials
@@ -850,8 +849,7 @@ class AwsProvider(Provider):
mfa_info = self.__input_role_mfa_token_and_code__()
assume_role_arguments["SerialNumber"] = mfa_info.arn
assume_role_arguments["TokenCode"] = mfa_info.totp
sts_client = create_sts_session(session, AWS_STS_GLOBAL_ENDPOINT_REGION)
sts_client = create_sts_session(session, assumed_role_info.sts_region)
assumed_credentials = sts_client.assume_role(**assume_role_arguments)
# Convert the UTC datetime object to your local timezone
credentials_expiration_local_time = (
@@ -998,7 +996,7 @@ def get_aws_region_for_sts(session_region: str, input_regions: set[str]) -> str:
# TODO: This can be moved to another class since it doesn't need self
def create_sts_session(
session: session.Session, aws_region: str
session: session.Session, aws_region: str = AWS_STS_GLOBAL_ENDPOINT_REGION
) -> session.Session.client:
"""
Create an STS session client.
+1
View File
@@ -1,3 +1,4 @@
AWS_STS_GLOBAL_ENDPOINT_REGION = "us-east-1"
AWS_REGION_US_EAST_1 = "us-east-1"
BOTO3_USER_AGENT_EXTRA = "APN_1826889"
ROLE_SESSION_NAME = "ProwlerAssessmentSession"
+1
View File
@@ -34,6 +34,7 @@ class AWSAssumeRoleInfo:
external_id: str
mfa_enabled: bool
role_session_name: str
sts_region: str
@dataclass
+85 -37
View File
@@ -41,6 +41,7 @@ from tests.providers.aws.utils import (
AWS_ACCOUNT_NUMBER,
AWS_CHINA_PARTITION,
AWS_COMMERCIAL_PARTITION,
AWS_GOV_CLOUD_ACCOUNT_ARN,
AWS_GOV_CLOUD_PARTITION,
AWS_ISO_PARTITION,
AWS_REGION_CN_NORTH_1,
@@ -436,6 +437,7 @@ class TestAWSProvider:
external_id=arguments.external_id,
mfa_enabled=True, # <- MFA configuration
role_session_name=arguments.role_session_name,
sts_region=AWS_REGION_US_EAST_1,
)
credentials = aws_provider._assumed_role_configuration.credentials
@@ -461,51 +463,97 @@ class TestAWSProvider:
arguments = Namespace()
arguments.mfa = False
role_name = "test-role"
arguments.role = f"arn:aws:iam::{AWS_ACCOUNT_NUMBER}:role/{role_name}"
arguments.role = (
f"arn:{AWS_COMMERCIAL_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/{role_name}"
)
arguments.session_duration = 900
arguments.role_session_name = "ProwlerAssessmentSession"
with patch(
"prowler.providers.aws.aws_provider.AwsProvider.__input_role_mfa_token_and_code__",
return_value=AWSMFAInfo(
arn=f"arn:aws:iam::{AWS_ACCOUNT_NUMBER}:mfa/test-role-mfa",
totp="111111",
),
):
aws_provider = AwsProvider(arguments)
assert (
aws_provider.session.current_session.region_name == AWS_REGION_US_EAST_1
)
assert aws_provider.identity.account == AWS_ACCOUNT_NUMBER
assert aws_provider.identity.account_arn == AWS_ACCOUNT_ARN
assert aws_provider.identity.partition == AWS_COMMERCIAL_PARTITION
assert isinstance(
aws_provider._assumed_role_configuration.info, AWSAssumeRoleInfo
)
assert aws_provider._assumed_role_configuration.info == AWSAssumeRoleInfo(
role_arn=ARN(arn=arguments.role),
session_duration=arguments.session_duration,
external_id=None,
mfa_enabled=False, # <- MFA configuration
role_session_name=arguments.role_session_name,
)
aws_provider = AwsProvider(arguments)
assert aws_provider.session.current_session.region_name == AWS_REGION_US_EAST_1
assert aws_provider.identity.account == AWS_ACCOUNT_NUMBER
assert aws_provider.identity.account_arn == AWS_ACCOUNT_ARN
assert aws_provider.identity.partition == AWS_COMMERCIAL_PARTITION
assert isinstance(
aws_provider._assumed_role_configuration.info, AWSAssumeRoleInfo
)
assert aws_provider._assumed_role_configuration.info == AWSAssumeRoleInfo(
role_arn=ARN(arn=arguments.role),
session_duration=arguments.session_duration,
external_id=None,
mfa_enabled=False, # <- MFA configuration
role_session_name=arguments.role_session_name,
sts_region=AWS_REGION_US_EAST_1,
)
credentials = aws_provider._assumed_role_configuration.credentials
assert isinstance(credentials, AWSCredentials)
credentials = aws_provider._assumed_role_configuration.credentials
assert isinstance(credentials, AWSCredentials)
assert credentials.aws_access_key_id
assert len(credentials.aws_access_key_id) == 20
assert search(r"^ASIA.*$", credentials.aws_access_key_id)
assert credentials.aws_access_key_id
assert len(credentials.aws_access_key_id) == 20
assert search(r"^ASIA.*$", credentials.aws_access_key_id)
assert credentials.aws_session_token
assert len(credentials.aws_session_token) == 356
assert search(r"^FQoGZXIvYXdzE.*$", credentials.aws_session_token)
assert credentials.aws_session_token
assert len(credentials.aws_session_token) == 356
assert search(r"^FQoGZXIvYXdzE.*$", credentials.aws_session_token)
assert credentials.aws_secret_access_key
assert len(credentials.aws_secret_access_key) == 40
assert credentials.aws_secret_access_key
assert len(credentials.aws_secret_access_key) == 40
assert credentials.expiration
# assert credentials.expiration == datetime.now(tzinfo=tzutc())
assert credentials.expiration
# assert credentials.expiration == datetime.now(tzinfo=tzutc())
@mock_aws
def test_aws_provider_assume_role_without_mfa_gov_cloud(self, monkeypatch):
# Set AWS_DEFAULT_REGION = 'us-gov-east-1' since is set by default to 'us-east-1
monkeypatch.setenv("AWS_DEFAULT_REGION", AWS_REGION_GOV_CLOUD_US_EAST_1)
# Variables
arguments = Namespace()
arguments.mfa = False
role_name = "test-role"
arguments.role = (
f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/{role_name}"
)
arguments.session_duration = 900
arguments.role_session_name = "ProwlerAssessmentSession"
aws_provider = AwsProvider(arguments)
assert (
aws_provider.session.current_session.region_name
== AWS_REGION_GOV_CLOUD_US_EAST_1
)
assert aws_provider.identity.account == AWS_ACCOUNT_NUMBER
assert aws_provider.identity.account_arn == AWS_GOV_CLOUD_ACCOUNT_ARN
assert aws_provider.identity.partition == AWS_GOV_CLOUD_PARTITION
assert isinstance(
aws_provider._assumed_role_configuration.info, AWSAssumeRoleInfo
)
assert aws_provider._assumed_role_configuration.info == AWSAssumeRoleInfo(
role_arn=ARN(arn=arguments.role),
session_duration=arguments.session_duration,
external_id=None,
mfa_enabled=False, # <- MFA configuration
role_session_name=arguments.role_session_name,
sts_region=AWS_REGION_GOV_CLOUD_US_EAST_1,
)
credentials = aws_provider._assumed_role_configuration.credentials
assert isinstance(credentials, AWSCredentials)
assert credentials.aws_access_key_id
assert len(credentials.aws_access_key_id) == 20
assert search(r"^ASIA.*$", credentials.aws_access_key_id)
assert credentials.aws_session_token
assert len(credentials.aws_session_token) == 356
assert search(r"^FQoGZXIvYXdzE.*$", credentials.aws_session_token)
assert credentials.aws_secret_access_key
assert len(credentials.aws_secret_access_key) == 40
assert credentials.expiration
# assert credentials.expiration == datetime.now(tzinfo=tzutc())
@mock_aws
def test_aws_provider_config(self):
+11 -8
View File
@@ -12,10 +12,19 @@ from prowler.config.config import (
from prowler.providers.aws.aws_provider import AwsProvider
from prowler.providers.common.models import Audit_Metadata
# AWS Partitions
AWS_COMMERCIAL_PARTITION = "aws"
AWS_GOV_CLOUD_PARTITION = "aws-us-gov"
AWS_CHINA_PARTITION = "aws-cn"
AWS_ISO_PARTITION = "aws-iso"
# Root AWS Account
AWS_ACCOUNT_NUMBER = "123456789012"
AWS_ACCOUNT_ARN = f"arn:aws:iam::{AWS_ACCOUNT_NUMBER}:root"
AWS_COMMERCIAL_PARTITION = "aws"
AWS_ACCOUNT_ARN = f"arn:{AWS_COMMERCIAL_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:root"
AWS_GOV_CLOUD_ACCOUNT_ARN = (
f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:root"
)
# Commercial Regions
AWS_REGION_US_EAST_1 = "us-east-1"
@@ -42,12 +51,6 @@ AWS_REGION_GOV_CLOUD_US_EAST_1 = "us-gov-east-1"
# Iso Regions
AWS_REGION_ISO_GLOBAL = "aws-iso-global"
# AWS Partitions
AWS_COMMERCIAL_PARTITION = "aws"
AWS_GOV_CLOUD_PARTITION = "aws-us-gov"
AWS_CHINA_PARTITION = "aws-cn"
AWS_ISO_PARTITION = "aws-iso"
# EC2
EXAMPLE_AMI_ID = "ami-12c6146b"