From bfc104795e0aae6c888e194f47aa153b11d46d7b Mon Sep 17 00:00:00 2001 From: Roja Reddy Sareddy Date: Thu, 6 Aug 2026 07:48:41 -0700 Subject: [PATCH 1/3] feat(evaluate): Add pre-validation for execution role permissions Validate execution role permissions before submitting evaluation jobs to fail fast with actionable guidance instead of failing during execution. Reuses existing _simulate_denied_actions and _role_trusts_service from iam_role_resolver.py. Applies to all evaluators via BaseEvaluator. --- .../common_utils/role_permission_validator.py | 195 +++++++++++++++++ .../train/evaluate/base_evaluator.py | 10 + .../train/evaluate/llm_as_judge_evaluator.py | 1 + .../evaluate/test_bedrock_role_validation.py | 199 ++++++++++++++++++ 4 files changed, 405 insertions(+) create mode 100644 sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py create mode 100644 sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py diff --git a/sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py b/sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py new file mode 100644 index 0000000000..3ddbc01fd1 --- /dev/null +++ b/sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py @@ -0,0 +1,195 @@ +"""Pre-validation for execution role permissions required by evaluation jobs. + +Validates that the execution role has the permissions and trust relationships +needed to run model evaluation jobs, providing actionable error messages at +submission time instead of failing 30+ minutes later during execution. + +All validation is best-effort: if the caller lacks iam:SimulatePrincipalPolicy +(common in Studio/notebooks), a warning is logged and execution proceeds. + +This module reuses the IAM simulation helpers from iam_role_resolver.py. +""" + +import logging +from typing import List, Optional + +from botocore.exceptions import ClientError + +from sagemaker.core.helper.iam_role_resolver import ( + _simulate_denied_actions, + _role_trusts_service, +) + +logger = logging.getLogger(__name__) + +_PREREQS_DOC_URL = ( + "https://docs.aws.amazon.com/sagemaker/latest/dg/" + "model-customize-open-weight-prereq.html" +) + +# Bedrock actions required on the execution role for evaluation jobs +_BEDROCK_EVAL_ACTIONS = [ + "bedrock:CreateEvaluationJob", + "bedrock:GetEvaluationJob", +] + +# MLflow actions required when experiment tracking is enabled +_MLFLOW_ACTIONS = [ + "sagemaker-mlflow:GetExperimentByName", + "sagemaker-mlflow:CreateExperiment", + "sagemaker-mlflow:CreateRun", + "sagemaker-mlflow:LogBatch", +] + + +def validate_evaluation_role_permissions( + role_arn: str, + sagemaker_session, + mlflow_enabled: bool = False, +) -> None: + """Validate that the execution role has permissions required for evaluation. + + Checks that the execution role can: + 1. Call bedrock:CreateEvaluationJob and bedrock:GetEvaluationJob. + 2. Be assumed by bedrock.amazonaws.com (trust relationship). + 3. (If mlflow_enabled) Call sagemaker-mlflow:* APIs for experiment tracking. + + Uses the same IAM simulation infrastructure as resolve_and_validate_role(). + Best-effort: warns and proceeds if caller cannot simulate. + + Args: + role_arn: The resolved execution role ARN. + sagemaker_session: SageMaker session (used to get boto session). + mlflow_enabled: Whether MLflow tracking is configured. + + Raises: + ValueError: If the role definitively lacks required permissions or trust. + """ + iam_client = _get_iam_client(sagemaker_session) + if iam_client is None: + return + + # 1. Check Bedrock permissions + _check_permissions( + iam_client, + role_arn, + _BEDROCK_EVAL_ACTIONS, + error_context=( + "Model evaluation requires these permissions on the execution role " + "to create and monitor Bedrock evaluation jobs." + ), + ) + + # 2. Check trust relationship for bedrock.amazonaws.com + _check_trust(iam_client, role_arn) + + # 3. Check MLflow permissions if tracking is enabled + if mlflow_enabled: + _check_permissions( + iam_client, + role_arn, + _MLFLOW_ACTIONS, + error_context=( + "MLflow experiment tracking is enabled (mlflow_resource_arn was provided), " + "but the execution role cannot access MLflow APIs." + ), + ) + + logger.info("Execution role '%s' validated for evaluation.", role_arn) + + +def _check_permissions( + iam_client, + role_arn: str, + actions: List[str], + error_context: str, +) -> None: + """Check that the role has the specified permissions via SimulatePrincipalPolicy.""" + try: + denied = _simulate_denied_actions(iam_client, role_arn, actions) + if denied: + raise ValueError( + f"IAM role '{role_arn}' is missing required permissions: " + f"{', '.join(denied)}. " + f"{error_context} " + f"To fix this, attach the 'AmazonSageMakerModelCustomizationCoreAccess' " + f"managed policy to your role, or add an inline policy granting " + f"the missing actions. " + f"See: {_PREREQS_DOC_URL}" + ) + except ClientError as e: + if _is_access_denied(e): + logger.warning( + "Could not verify permissions for role '%s' (caller lacks " + "iam:SimulatePrincipalPolicy). If the job fails with " + "AccessDeniedException, ensure the role has: %s", + role_arn, + ", ".join(actions), + ) + elif _is_no_such_entity(e): + pass # Role gone; let it fail downstream with clearer error + else: + raise + except ValueError: + raise + except Exception as e: + logger.info("Permission check failed unexpectedly: %s; skipping.", e) + + +def _check_trust(iam_client, role_arn: str) -> None: + """Check that the role trusts bedrock.amazonaws.com.""" + try: + trusts_bedrock = _role_trusts_service(iam_client, role_arn, "bedrock") + if trusts_bedrock is False: + raise ValueError( + f"IAM role '{role_arn}' trust policy does not include " + f"'bedrock.amazonaws.com'. Model evaluation requires " + f"Bedrock to assume your execution role to run the evaluation job. " + f"Add the following trust statement to your role:\n" + f'{{"Effect": "Allow", "Principal": {{"Service": ' + f'"bedrock.amazonaws.com"}}, "Action": "sts:AssumeRole"}}\n' + f"See: {_PREREQS_DOC_URL}" + ) + except ClientError as e: + if _is_access_denied(e): + logger.warning( + "Could not verify trust policy for role '%s'. If the evaluation " + "fails with 'Could not assume role', ensure bedrock.amazonaws.com " + "is in the role's trust policy.", + role_arn, + ) + elif _is_no_such_entity(e): + pass + else: + raise + except ValueError: + raise + except Exception as e: + logger.info("Trust check failed unexpectedly: %s; skipping.", e) + + +# --------------------------------------------------------------------------- +# Internal helpers +# --------------------------------------------------------------------------- + +def _get_iam_client(sagemaker_session): + """Get an IAM client from the session, or None if unavailable.""" + import boto3 + + try: + if sagemaker_session and hasattr(sagemaker_session, 'boto_session'): + return sagemaker_session.boto_session.client("iam") + return boto3.Session().client("iam") + except Exception: + logger.info("Could not create IAM client for role validation; skipping.") + return None + + +def _is_access_denied(error) -> bool: + code = error.response.get("Error", {}).get("Code", "") + return code in ("AccessDenied", "AccessDeniedException") + + +def _is_no_such_entity(error) -> bool: + code = error.response.get("Error", {}).get("Code", "") + return code in ("NoSuchEntity", "NoSuchEntityException") diff --git a/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py b/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py index 29ab44133a..48ee00d3ca 100644 --- a/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py +++ b/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py @@ -30,6 +30,9 @@ _is_nova_model, ) from sagemaker.train.common_utils.recipe_utils import resolve_recipe, get_resolved_recipe_from_context +from sagemaker.train.common_utils.role_permission_validator import ( + validate_evaluation_role_permissions, +) from sagemaker.train.common_utils.validator import validate_hyperpod_compute from sagemaker.train.defaults import TrainDefaults from sagemaker.train.recipe_resolver import flatten_resolved_recipe @@ -748,6 +751,13 @@ def _get_aws_execution_context(self) -> Dict[str, str]: # Extract account ID from role ARN account_id = role_arn.split(':')[4] if ':' in role_arn else '052150106756' + + # Validate execution role has permissions required for evaluation + validate_evaluation_role_permissions( + role_arn, + self.sagemaker_session, + mlflow_enabled=bool(getattr(self, 'mlflow_resource_arn', None)), + ) return { 'role_arn': role_arn, diff --git a/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py b/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py index a157c05701..edc39ff45f 100644 --- a/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py +++ b/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py @@ -752,6 +752,7 @@ def _get_llmaj_template_additions(self, eval_name: str) -> dict: 'evaluate_base_model': self.evaluate_base_model, } + @_telemetry_emitter( feature=Feature.MODEL_CUSTOMIZATION, func_name="LLMAsJudgeEvaluator.evaluate", diff --git a/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py b/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py new file mode 100644 index 0000000000..15b4968c6f --- /dev/null +++ b/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py @@ -0,0 +1,199 @@ +"""Unit tests for execution role permission validation for evaluation jobs.""" + +import pytest +from unittest.mock import MagicMock, patch +from botocore.exceptions import ClientError + +from sagemaker.train.common_utils.role_permission_validator import ( + validate_evaluation_role_permissions, +) + + +class TestValidateEvaluationRolePermissions: + """Tests for validate_evaluation_role_permissions.""" + + def _make_session(self, iam_client): + session = MagicMock() + session.boto_session.client.return_value = iam_client + return session + + def _mock_simulate_all_allowed(self, iam_client): + """All actions return allowed.""" + paginator = MagicMock() + paginator.paginate.return_value = [ + { + "EvaluationResults": [ + {"EvalActionName": action, "EvalDecision": "allowed"} + for action in [ + "bedrock:CreateEvaluationJob", + "bedrock:GetEvaluationJob", + "sagemaker-mlflow:GetExperimentByName", + "sagemaker-mlflow:CreateExperiment", + "sagemaker-mlflow:CreateRun", + "sagemaker-mlflow:LogBatch", + ] + ] + } + ] + iam_client.get_paginator.return_value = paginator + + def _mock_simulate_with_denied(self, iam_client, denied_actions): + """Return denied for specified actions, allowed for the rest.""" + def paginate_side_effect(**kwargs): + results = [ + {"EvalActionName": a, "EvalDecision": "implicitDeny" if a in denied_actions else "allowed"} + for a in kwargs["ActionNames"] + ] + return [{"EvaluationResults": results}] + + paginator = MagicMock() + paginator.paginate.side_effect = paginate_side_effect + iam_client.get_paginator.return_value = paginator + + def _mock_role_trusts_bedrock(self, iam_client, trusts=True): + """Mock _role_trusts_service result via get_role.""" + services = ["sagemaker.amazonaws.com"] + if trusts: + services.append("bedrock.amazonaws.com") + iam_client.get_role.return_value = { + "Role": { + "AssumeRolePolicyDocument": { + "Version": "2012-10-17", + "Statement": [{ + "Effect": "Allow", + "Principal": {"Service": services}, + "Action": "sts:AssumeRole", + }], + } + } + } + + # --- Bedrock permission tests --- + + def test_passes_when_all_valid(self): + iam_client = MagicMock() + self._mock_simulate_all_allowed(iam_client) + self._mock_role_trusts_bedrock(iam_client, trusts=True) + session = self._make_session(iam_client) + + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + def test_raises_when_create_evaluation_job_denied(self): + iam_client = MagicMock() + self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob"]) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="bedrock:CreateEvaluationJob"): + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + def test_raises_when_get_evaluation_job_denied(self): + iam_client = MagicMock() + self._mock_simulate_with_denied(iam_client, ["bedrock:GetEvaluationJob"]) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="bedrock:GetEvaluationJob"): + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + def test_raises_when_both_bedrock_actions_denied(self): + iam_client = MagicMock() + self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob", "bedrock:GetEvaluationJob"]) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="bedrock:CreateEvaluationJob"): + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + # --- Trust policy tests --- + + def test_raises_when_trust_missing_bedrock(self): + iam_client = MagicMock() + self._mock_simulate_all_allowed(iam_client) + self._mock_role_trusts_bedrock(iam_client, trusts=False) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="bedrock.amazonaws.com"): + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + # --- MLflow permission tests --- + + def test_raises_when_mlflow_denied_and_enabled(self): + iam_client = MagicMock() + self._mock_simulate_with_denied(iam_client, ["sagemaker-mlflow:GetExperimentByName"]) + self._mock_role_trusts_bedrock(iam_client, trusts=True) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="sagemaker-mlflow:GetExperimentByName"): + validate_evaluation_role_permissions( + "arn:aws:iam::123456789012:role/MyRole", session, mlflow_enabled=True + ) + + def test_skips_mlflow_check_when_not_enabled(self): + iam_client = MagicMock() + # Bedrock allowed, but MLflow would be denied - doesn't matter since not enabled + self._mock_simulate_with_denied(iam_client, ["sagemaker-mlflow:GetExperimentByName"]) + self._mock_role_trusts_bedrock(iam_client, trusts=True) + session = self._make_session(iam_client) + + # Should not raise since mlflow_enabled=False (default) + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + # --- Graceful degradation tests --- + + def test_warns_when_simulate_access_denied(self): + iam_client = MagicMock() + paginator = MagicMock() + paginator.paginate.side_effect = ClientError( + {"Error": {"Code": "AccessDenied", "Message": ""}}, "SimulatePrincipalPolicy" + ) + iam_client.get_paginator.return_value = paginator + # Trust check also needs to handle gracefully + iam_client.get_role.side_effect = ClientError( + {"Error": {"Code": "AccessDenied", "Message": ""}}, "GetRole" + ) + session = self._make_session(iam_client) + + # Should not raise + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + def test_warns_when_get_role_access_denied(self): + iam_client = MagicMock() + self._mock_simulate_all_allowed(iam_client) + iam_client.get_role.side_effect = ClientError( + {"Error": {"Code": "AccessDenied", "Message": ""}}, "GetRole" + ) + session = self._make_session(iam_client) + + # Should not raise + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + def test_skips_when_session_is_none(self): + with patch("boto3.Session", side_effect=Exception("no creds")): + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", None) + + # --- Error message quality tests --- + + def test_error_includes_doc_link(self): + iam_client = MagicMock() + self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob"]) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="model-customize-open-weight-prereq"): + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + def test_error_suggests_managed_policy(self): + iam_client = MagicMock() + self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob"]) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="AmazonSageMakerModelCustomizationCoreAccess"): + validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + + def test_mlflow_error_mentions_tracking_enabled(self): + iam_client = MagicMock() + self._mock_simulate_with_denied(iam_client, ["sagemaker-mlflow:CreateRun"]) + self._mock_role_trusts_bedrock(iam_client, trusts=True) + session = self._make_session(iam_client) + + with pytest.raises(ValueError, match="MLflow experiment tracking is enabled"): + validate_evaluation_role_permissions( + "arn:aws:iam::123456789012:role/MyRole", session, mlflow_enabled=True + ) From da692d24227a73f4a65daeef58d03f02c8b4fd88 Mon Sep 17 00:00:00 2001 From: Roja Reddy Sareddy Date: Thu, 6 Aug 2026 08:23:39 -0700 Subject: [PATCH 2/3] refactor(evaluate): Use evaluation role type via iam_role_resolver Add 'evaluation' role type to IAM_POLICY_CONFIG with Bedrock and MLflow permissions. BaseEvaluator now calls resolve_and_validate_role with role_type='evaluation' which validates permissions and trust using the existing iam_role_resolver infrastructure. Removes the standalone role_permission_validator utility in favor of the existing centralized approach. --- .../src/sagemaker/core/helper/iam_policies.py | 145 ++++++++ .../core/helper/iam_role_resolver.py | 2 +- .../common_utils/role_permission_validator.py | 195 ----------- .../train/evaluate/base_evaluator.py | 14 +- .../evaluate/test_bedrock_role_validation.py | 326 +++++++++--------- 5 files changed, 311 insertions(+), 371 deletions(-) delete mode 100644 sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py diff --git a/sagemaker-core/src/sagemaker/core/helper/iam_policies.py b/sagemaker-core/src/sagemaker/core/helper/iam_policies.py index f583b3b586..4b868bc9f3 100644 --- a/sagemaker-core/src/sagemaker/core/helper/iam_policies.py +++ b/sagemaker-core/src/sagemaker/core/helper/iam_policies.py @@ -812,4 +812,149 @@ }, }, }, + "evaluation": { + "role_name": "SageMaker-AutoRole-Evaluation", + "trust_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Principal": { + "Service": [ + "sagemaker.amazonaws.com", + "bedrock.amazonaws.com", + ] + }, + "Action": "sts:AssumeRole", + "Condition": { + "StringEquals": {"aws:SourceAccount": "ACCOUNT_PLACEHOLDER"} + }, + } + ], + }, + "policies": { + # --- Training permissions (superset of "training" role type) --- + "s3_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": [ + "s3:GetObject", + "s3:PutObject", + "s3:ListBucket", + "s3:GetBucketLocation", + ], + "Resource": "S3_PLACEHOLDER", + } + ], + }, + "ecr_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": ["ecr:GetAuthorizationToken"], + "Resource": "*", + }, + { + "Effect": "Allow", + "Action": [ + "ecr:GetDownloadUrlForLayer", + "ecr:BatchGetImage", + "ecr:BatchCheckLayerAvailability", + ], + "Resource": "arn:aws:ecr:*:*:repository/*", + }, + ], + }, + "cloudwatch_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": ["cloudwatch:PutMetricData"], + "Resource": "*", + }, + { + "Effect": "Allow", + "Action": [ + "logs:CreateLogGroup", + "logs:CreateLogStream", + "logs:PutLogEvents", + "logs:DescribeLogStreams", + ], + "Resource": "arn:aws:logs:*:*:log-group:/aws/sagemaker/TrainingJobs*", + }, + ], + }, + "kms_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": ["kms:Encrypt", "kms:Decrypt", "kms:GenerateDataKey"], + "Resource": "KMS_PLACEHOLDER", + } + ], + }, + # --- Evaluation-specific permissions --- + "bedrock_evaluation_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": [ + "bedrock:CreateEvaluationJob", + "bedrock:GetEvaluationJob", + "bedrock:InvokeModel", + "bedrock:InvokeModelWithResponseStream", + ], + "Resource": "*", + } + ], + }, + "mlflow_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": [ + "sagemaker-mlflow:GetExperimentByName", + "sagemaker-mlflow:CreateExperiment", + "sagemaker-mlflow:CreateRun", + "sagemaker-mlflow:LogBatch", + "sagemaker-mlflow:LogMetric", + "sagemaker-mlflow:LogParam", + "sagemaker-mlflow:SetTag", + "sagemaker-mlflow:UpdateRun", + ], + "Resource": "arn:aws:sagemaker:*:*:mlflow-app/*", + } + ], + }, + "sagemaker_evaluation_policy": { + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": [ + "sagemaker:CreateTrainingJob", + "sagemaker:DescribeTrainingJob", + "sagemaker:StopTrainingJob", + "sagemaker:CreatePipeline", + "sagemaker:DescribePipeline", + "sagemaker:StartPipelineExecution", + "sagemaker:DescribePipelineExecution", + "sagemaker:AddTags", + ], + "Resource": [ + "arn:aws:sagemaker:*:*:training-job/*", + "arn:aws:sagemaker:*:*:pipeline/*", + ], + } + ], + }, + }, + }, } diff --git a/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py b/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py index 9abf33fab6..c5f1a57cd1 100644 --- a/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py +++ b/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py @@ -25,7 +25,7 @@ logger = logging.getLogger(__name__) -ROLE_TYPES = ("training", "serving", "pipeline", "feature_store", "bedrock", "hyperpod") +ROLE_TYPES = ("training", "serving", "pipeline", "feature_store", "bedrock", "hyperpod", "evaluation") # Permissions the HyperPod CLI flow needs on the *caller* identity — the local # principal that runs `hyperpod connect-cluster` and `hyperpod start-job`. The CLI diff --git a/sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py b/sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py deleted file mode 100644 index 3ddbc01fd1..0000000000 --- a/sagemaker-train/src/sagemaker/train/common_utils/role_permission_validator.py +++ /dev/null @@ -1,195 +0,0 @@ -"""Pre-validation for execution role permissions required by evaluation jobs. - -Validates that the execution role has the permissions and trust relationships -needed to run model evaluation jobs, providing actionable error messages at -submission time instead of failing 30+ minutes later during execution. - -All validation is best-effort: if the caller lacks iam:SimulatePrincipalPolicy -(common in Studio/notebooks), a warning is logged and execution proceeds. - -This module reuses the IAM simulation helpers from iam_role_resolver.py. -""" - -import logging -from typing import List, Optional - -from botocore.exceptions import ClientError - -from sagemaker.core.helper.iam_role_resolver import ( - _simulate_denied_actions, - _role_trusts_service, -) - -logger = logging.getLogger(__name__) - -_PREREQS_DOC_URL = ( - "https://docs.aws.amazon.com/sagemaker/latest/dg/" - "model-customize-open-weight-prereq.html" -) - -# Bedrock actions required on the execution role for evaluation jobs -_BEDROCK_EVAL_ACTIONS = [ - "bedrock:CreateEvaluationJob", - "bedrock:GetEvaluationJob", -] - -# MLflow actions required when experiment tracking is enabled -_MLFLOW_ACTIONS = [ - "sagemaker-mlflow:GetExperimentByName", - "sagemaker-mlflow:CreateExperiment", - "sagemaker-mlflow:CreateRun", - "sagemaker-mlflow:LogBatch", -] - - -def validate_evaluation_role_permissions( - role_arn: str, - sagemaker_session, - mlflow_enabled: bool = False, -) -> None: - """Validate that the execution role has permissions required for evaluation. - - Checks that the execution role can: - 1. Call bedrock:CreateEvaluationJob and bedrock:GetEvaluationJob. - 2. Be assumed by bedrock.amazonaws.com (trust relationship). - 3. (If mlflow_enabled) Call sagemaker-mlflow:* APIs for experiment tracking. - - Uses the same IAM simulation infrastructure as resolve_and_validate_role(). - Best-effort: warns and proceeds if caller cannot simulate. - - Args: - role_arn: The resolved execution role ARN. - sagemaker_session: SageMaker session (used to get boto session). - mlflow_enabled: Whether MLflow tracking is configured. - - Raises: - ValueError: If the role definitively lacks required permissions or trust. - """ - iam_client = _get_iam_client(sagemaker_session) - if iam_client is None: - return - - # 1. Check Bedrock permissions - _check_permissions( - iam_client, - role_arn, - _BEDROCK_EVAL_ACTIONS, - error_context=( - "Model evaluation requires these permissions on the execution role " - "to create and monitor Bedrock evaluation jobs." - ), - ) - - # 2. Check trust relationship for bedrock.amazonaws.com - _check_trust(iam_client, role_arn) - - # 3. Check MLflow permissions if tracking is enabled - if mlflow_enabled: - _check_permissions( - iam_client, - role_arn, - _MLFLOW_ACTIONS, - error_context=( - "MLflow experiment tracking is enabled (mlflow_resource_arn was provided), " - "but the execution role cannot access MLflow APIs." - ), - ) - - logger.info("Execution role '%s' validated for evaluation.", role_arn) - - -def _check_permissions( - iam_client, - role_arn: str, - actions: List[str], - error_context: str, -) -> None: - """Check that the role has the specified permissions via SimulatePrincipalPolicy.""" - try: - denied = _simulate_denied_actions(iam_client, role_arn, actions) - if denied: - raise ValueError( - f"IAM role '{role_arn}' is missing required permissions: " - f"{', '.join(denied)}. " - f"{error_context} " - f"To fix this, attach the 'AmazonSageMakerModelCustomizationCoreAccess' " - f"managed policy to your role, or add an inline policy granting " - f"the missing actions. " - f"See: {_PREREQS_DOC_URL}" - ) - except ClientError as e: - if _is_access_denied(e): - logger.warning( - "Could not verify permissions for role '%s' (caller lacks " - "iam:SimulatePrincipalPolicy). If the job fails with " - "AccessDeniedException, ensure the role has: %s", - role_arn, - ", ".join(actions), - ) - elif _is_no_such_entity(e): - pass # Role gone; let it fail downstream with clearer error - else: - raise - except ValueError: - raise - except Exception as e: - logger.info("Permission check failed unexpectedly: %s; skipping.", e) - - -def _check_trust(iam_client, role_arn: str) -> None: - """Check that the role trusts bedrock.amazonaws.com.""" - try: - trusts_bedrock = _role_trusts_service(iam_client, role_arn, "bedrock") - if trusts_bedrock is False: - raise ValueError( - f"IAM role '{role_arn}' trust policy does not include " - f"'bedrock.amazonaws.com'. Model evaluation requires " - f"Bedrock to assume your execution role to run the evaluation job. " - f"Add the following trust statement to your role:\n" - f'{{"Effect": "Allow", "Principal": {{"Service": ' - f'"bedrock.amazonaws.com"}}, "Action": "sts:AssumeRole"}}\n' - f"See: {_PREREQS_DOC_URL}" - ) - except ClientError as e: - if _is_access_denied(e): - logger.warning( - "Could not verify trust policy for role '%s'. If the evaluation " - "fails with 'Could not assume role', ensure bedrock.amazonaws.com " - "is in the role's trust policy.", - role_arn, - ) - elif _is_no_such_entity(e): - pass - else: - raise - except ValueError: - raise - except Exception as e: - logger.info("Trust check failed unexpectedly: %s; skipping.", e) - - -# --------------------------------------------------------------------------- -# Internal helpers -# --------------------------------------------------------------------------- - -def _get_iam_client(sagemaker_session): - """Get an IAM client from the session, or None if unavailable.""" - import boto3 - - try: - if sagemaker_session and hasattr(sagemaker_session, 'boto_session'): - return sagemaker_session.boto_session.client("iam") - return boto3.Session().client("iam") - except Exception: - logger.info("Could not create IAM client for role validation; skipping.") - return None - - -def _is_access_denied(error) -> bool: - code = error.response.get("Error", {}).get("Code", "") - return code in ("AccessDenied", "AccessDeniedException") - - -def _is_no_such_entity(error) -> bool: - code = error.response.get("Error", {}).get("Code", "") - return code in ("NoSuchEntity", "NoSuchEntityException") diff --git a/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py b/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py index 48ee00d3ca..c44a1a31f4 100644 --- a/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py +++ b/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py @@ -30,9 +30,6 @@ _is_nova_model, ) from sagemaker.train.common_utils.recipe_utils import resolve_recipe, get_resolved_recipe_from_context -from sagemaker.train.common_utils.role_permission_validator import ( - validate_evaluation_role_permissions, -) from sagemaker.train.common_utils.validator import validate_hyperpod_compute from sagemaker.train.defaults import TrainDefaults from sagemaker.train.recipe_resolver import flatten_resolved_recipe @@ -737,10 +734,10 @@ def _get_aws_execution_context(self) -> Dict[str, str]: # This is the job execution role for the # serverless / SMTJ evaluation backends. The HyperPod backend submits via # the CLI under the caller's own credentials (see _submit_hyperpod_eval_job) - # and does not resolve a role here, so "training" is always correct here. + # and does not resolve a role here, so "evaluation" is correct here. role_arn = resolve_and_validate_role( provided_role=self.role, - role_type="training", + role_type="evaluation", sagemaker_session=self.sagemaker_session, ) @@ -751,13 +748,6 @@ def _get_aws_execution_context(self) -> Dict[str, str]: # Extract account ID from role ARN account_id = role_arn.split(':')[4] if ':' in role_arn else '052150106756' - - # Validate execution role has permissions required for evaluation - validate_evaluation_role_permissions( - role_arn, - self.sagemaker_session, - mlflow_enabled=bool(getattr(self, 'mlflow_resource_arn', None)), - ) return { 'role_arn': role_arn, diff --git a/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py b/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py index 15b4968c6f..1a698bbb34 100644 --- a/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py +++ b/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py @@ -1,199 +1,199 @@ -"""Unit tests for execution role permission validation for evaluation jobs.""" +"""Unit tests for the 'evaluation' role type in iam_role_resolver.""" import pytest from unittest.mock import MagicMock, patch from botocore.exceptions import ClientError -from sagemaker.train.common_utils.role_permission_validator import ( - validate_evaluation_role_permissions, +from sagemaker.core.helper.iam_role_resolver import ( + resolve_and_validate_role, + RoleValidationError, + _evaluate_permissions, + _role_trusts_service, + _get_smoke_test_actions, + _expected_trust_services, ) -class TestValidateEvaluationRolePermissions: - """Tests for validate_evaluation_role_permissions.""" +class TestEvaluationRoleType: + """Tests for the 'evaluation' role type configuration.""" - def _make_session(self, iam_client): - session = MagicMock() - session.boto_session.client.return_value = iam_client - return session + def test_evaluation_role_type_exists(self): + """The 'evaluation' role type should be recognized.""" + from sagemaker.core.helper.iam_policies import IAM_POLICY_CONFIG + assert "evaluation" in IAM_POLICY_CONFIG - def _mock_simulate_all_allowed(self, iam_client): - """All actions return allowed.""" + def test_evaluation_trust_includes_bedrock(self): + """Evaluation role type should require bedrock.amazonaws.com trust.""" + expected = _expected_trust_services("evaluation") + assert "bedrock.amazonaws.com" in expected + + def test_evaluation_trust_includes_sagemaker(self): + """Evaluation role type should require sagemaker.amazonaws.com trust.""" + expected = _expected_trust_services("evaluation") + assert "sagemaker.amazonaws.com" in expected + + def test_evaluation_smoke_actions_include_bedrock(self): + """Smoke test actions should include Bedrock evaluation actions.""" + actions = _get_smoke_test_actions("evaluation") + assert "bedrock:CreateEvaluationJob" in actions + assert "bedrock:GetEvaluationJob" in actions + + def test_evaluation_smoke_actions_include_bedrock_invoke(self): + """Smoke test actions should include Bedrock invoke actions.""" + actions = _get_smoke_test_actions("evaluation") + assert "bedrock:InvokeModel" in actions + + def test_resolve_raises_when_bedrock_permissions_denied(self): + """Should raise RoleValidationError when Bedrock permissions are denied.""" + mock_iam = MagicMock() + + # Simulate: bedrock:CreateEvaluationJob denied paginator = MagicMock() paginator.paginate.return_value = [ { "EvaluationResults": [ - {"EvalActionName": action, "EvalDecision": "allowed"} - for action in [ - "bedrock:CreateEvaluationJob", - "bedrock:GetEvaluationJob", - "sagemaker-mlflow:GetExperimentByName", - "sagemaker-mlflow:CreateExperiment", - "sagemaker-mlflow:CreateRun", - "sagemaker-mlflow:LogBatch", - ] + {"EvalActionName": "bedrock:CreateEvaluationJob", "EvalDecision": "implicitDeny"}, + {"EvalActionName": "bedrock:GetEvaluationJob", "EvalDecision": "allowed"}, + {"EvalActionName": "bedrock:InvokeModel", "EvalDecision": "allowed"}, + {"EvalActionName": "bedrock:InvokeModelWithResponseStream", "EvalDecision": "allowed"}, ] } ] - iam_client.get_paginator.return_value = paginator - - def _mock_simulate_with_denied(self, iam_client, denied_actions): - """Return denied for specified actions, allowed for the rest.""" - def paginate_side_effect(**kwargs): - results = [ - {"EvalActionName": a, "EvalDecision": "implicitDeny" if a in denied_actions else "allowed"} - for a in kwargs["ActionNames"] - ] - return [{"EvaluationResults": results}] + mock_iam.get_paginator.return_value = paginator + verdict, denied = _evaluate_permissions(mock_iam, "arn:aws:iam::123456789012:role/MyRole", "evaluation") + assert verdict is False + assert "bedrock:CreateEvaluationJob" in denied + + def test_resolve_passes_when_all_allowed(self): + """Should pass when all evaluation permissions are allowed.""" + mock_iam = MagicMock() + + actions = _get_smoke_test_actions("evaluation") paginator = MagicMock() - paginator.paginate.side_effect = paginate_side_effect - iam_client.get_paginator.return_value = paginator - - def _mock_role_trusts_bedrock(self, iam_client, trusts=True): - """Mock _role_trusts_service result via get_role.""" - services = ["sagemaker.amazonaws.com"] - if trusts: - services.append("bedrock.amazonaws.com") - iam_client.get_role.return_value = { + paginator.paginate.return_value = [ + { + "EvaluationResults": [ + {"EvalActionName": a, "EvalDecision": "allowed"} for a in actions + ] + } + ] + mock_iam.get_paginator.return_value = paginator + + verdict, denied = _evaluate_permissions(mock_iam, "arn:aws:iam::123456789012:role/MyRole", "evaluation") + assert verdict is True + assert denied == [] + + def test_trust_check_fails_without_bedrock(self): + """Should return False when trust policy lacks bedrock.amazonaws.com.""" + mock_iam = MagicMock() + mock_iam.get_role.return_value = { "Role": { "AssumeRolePolicyDocument": { "Version": "2012-10-17", "Statement": [{ "Effect": "Allow", - "Principal": {"Service": services}, + "Principal": {"Service": "sagemaker.amazonaws.com"}, "Action": "sts:AssumeRole", }], } } } - # --- Bedrock permission tests --- - - def test_passes_when_all_valid(self): - iam_client = MagicMock() - self._mock_simulate_all_allowed(iam_client) - self._mock_role_trusts_bedrock(iam_client, trusts=True) - session = self._make_session(iam_client) - - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - def test_raises_when_create_evaluation_job_denied(self): - iam_client = MagicMock() - self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob"]) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="bedrock:CreateEvaluationJob"): - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - def test_raises_when_get_evaluation_job_denied(self): - iam_client = MagicMock() - self._mock_simulate_with_denied(iam_client, ["bedrock:GetEvaluationJob"]) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="bedrock:GetEvaluationJob"): - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - def test_raises_when_both_bedrock_actions_denied(self): - iam_client = MagicMock() - self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob", "bedrock:GetEvaluationJob"]) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="bedrock:CreateEvaluationJob"): - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + result = _role_trusts_service(mock_iam, "arn:aws:iam::123456789012:role/MyRole", "evaluation") + assert result is False - # --- Trust policy tests --- - - def test_raises_when_trust_missing_bedrock(self): - iam_client = MagicMock() - self._mock_simulate_all_allowed(iam_client) - self._mock_role_trusts_bedrock(iam_client, trusts=False) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="bedrock.amazonaws.com"): - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - # --- MLflow permission tests --- - - def test_raises_when_mlflow_denied_and_enabled(self): - iam_client = MagicMock() - self._mock_simulate_with_denied(iam_client, ["sagemaker-mlflow:GetExperimentByName"]) - self._mock_role_trusts_bedrock(iam_client, trusts=True) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="sagemaker-mlflow:GetExperimentByName"): - validate_evaluation_role_permissions( - "arn:aws:iam::123456789012:role/MyRole", session, mlflow_enabled=True - ) + def test_trust_check_passes_with_bedrock(self): + """Should return True when trust policy includes both services.""" + mock_iam = MagicMock() + mock_iam.get_role.return_value = { + "Role": { + "AssumeRolePolicyDocument": { + "Version": "2012-10-17", + "Statement": [{ + "Effect": "Allow", + "Principal": {"Service": ["sagemaker.amazonaws.com", "bedrock.amazonaws.com"]}, + "Action": "sts:AssumeRole", + }], + } + } + } - def test_skips_mlflow_check_when_not_enabled(self): - iam_client = MagicMock() - # Bedrock allowed, but MLflow would be denied - doesn't matter since not enabled - self._mock_simulate_with_denied(iam_client, ["sagemaker-mlflow:GetExperimentByName"]) - self._mock_role_trusts_bedrock(iam_client, trusts=True) - session = self._make_session(iam_client) + result = _role_trusts_service(mock_iam, "arn:aws:iam::123456789012:role/MyRole", "evaluation") + assert result is True + + def test_resolve_and_validate_raises_on_trust_failure(self): + """Full resolve_and_validate_role should raise when trust fails.""" + with patch("sagemaker.core.helper.iam_role_resolver._get_boto_session") as mock_session: + mock_boto = MagicMock() + mock_session.return_value = mock_boto + + mock_iam = MagicMock() + mock_boto.client.return_value = mock_iam + + # Role exists + mock_iam.get_role.return_value = { + "Role": { + "Arn": "arn:aws:iam::123456789012:role/MyRole", + "AssumeRolePolicyDocument": { + "Version": "2012-10-17", + "Statement": [{ + "Effect": "Allow", + "Principal": {"Service": "sagemaker.amazonaws.com"}, + "Action": "sts:AssumeRole", + }], + } + } + } - # Should not raise since mlflow_enabled=False (default) - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) + # Permissions all allowed + actions = _get_smoke_test_actions("evaluation") + paginator = MagicMock() + paginator.paginate.return_value = [ + {"EvaluationResults": [{"EvalActionName": a, "EvalDecision": "allowed"} for a in actions]} + ] + mock_iam.get_paginator.return_value = paginator + + with pytest.raises(RoleValidationError, match="bedrock.amazonaws.com"): + resolve_and_validate_role( + provided_role="arn:aws:iam::123456789012:role/MyRole", + role_type="evaluation", + ) + + def test_resolve_and_validate_passes_with_correct_role(self): + """Full resolve_and_validate_role should pass with correct permissions and trust.""" + with patch("sagemaker.core.helper.iam_role_resolver._get_boto_session") as mock_session: + mock_boto = MagicMock() + mock_session.return_value = mock_boto + + mock_iam = MagicMock() + mock_boto.client.return_value = mock_iam + + # Role exists with correct trust + mock_iam.get_role.return_value = { + "Role": { + "Arn": "arn:aws:iam::123456789012:role/MyRole", + "AssumeRolePolicyDocument": { + "Version": "2012-10-17", + "Statement": [{ + "Effect": "Allow", + "Principal": {"Service": ["sagemaker.amazonaws.com", "bedrock.amazonaws.com"]}, + "Action": "sts:AssumeRole", + }], + } + } + } - # --- Graceful degradation tests --- + # All permissions allowed + actions = _get_smoke_test_actions("evaluation") + paginator = MagicMock() + paginator.paginate.return_value = [ + {"EvaluationResults": [{"EvalActionName": a, "EvalDecision": "allowed"} for a in actions]} + ] + mock_iam.get_paginator.return_value = paginator - def test_warns_when_simulate_access_denied(self): - iam_client = MagicMock() - paginator = MagicMock() - paginator.paginate.side_effect = ClientError( - {"Error": {"Code": "AccessDenied", "Message": ""}}, "SimulatePrincipalPolicy" - ) - iam_client.get_paginator.return_value = paginator - # Trust check also needs to handle gracefully - iam_client.get_role.side_effect = ClientError( - {"Error": {"Code": "AccessDenied", "Message": ""}}, "GetRole" - ) - session = self._make_session(iam_client) - - # Should not raise - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - def test_warns_when_get_role_access_denied(self): - iam_client = MagicMock() - self._mock_simulate_all_allowed(iam_client) - iam_client.get_role.side_effect = ClientError( - {"Error": {"Code": "AccessDenied", "Message": ""}}, "GetRole" - ) - session = self._make_session(iam_client) - - # Should not raise - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - def test_skips_when_session_is_none(self): - with patch("boto3.Session", side_effect=Exception("no creds")): - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", None) - - # --- Error message quality tests --- - - def test_error_includes_doc_link(self): - iam_client = MagicMock() - self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob"]) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="model-customize-open-weight-prereq"): - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - def test_error_suggests_managed_policy(self): - iam_client = MagicMock() - self._mock_simulate_with_denied(iam_client, ["bedrock:CreateEvaluationJob"]) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="AmazonSageMakerModelCustomizationCoreAccess"): - validate_evaluation_role_permissions("arn:aws:iam::123456789012:role/MyRole", session) - - def test_mlflow_error_mentions_tracking_enabled(self): - iam_client = MagicMock() - self._mock_simulate_with_denied(iam_client, ["sagemaker-mlflow:CreateRun"]) - self._mock_role_trusts_bedrock(iam_client, trusts=True) - session = self._make_session(iam_client) - - with pytest.raises(ValueError, match="MLflow experiment tracking is enabled"): - validate_evaluation_role_permissions( - "arn:aws:iam::123456789012:role/MyRole", session, mlflow_enabled=True + result = resolve_and_validate_role( + provided_role="arn:aws:iam::123456789012:role/MyRole", + role_type="evaluation", ) + assert result == "arn:aws:iam::123456789012:role/MyRole" From 3b9de1112375f66000f7e58521c0c042323bd303 Mon Sep 17 00:00:00 2001 From: Roja Reddy Sareddy Date: Fri, 7 Aug 2026 06:41:57 -0700 Subject: [PATCH 3/3] fix(evaluate): Address review feedback on role validation - Remove bedrock.amazonaws.com from evaluation trust policy. The serverless backend calls Bedrock using the role's own credentials; Bedrock does not assume the role, so trust is not required. - Update existing base_evaluator tests to expect role_type='evaluation'. - Remove stray blank line in llm_as_judge_evaluator.py. - Update test assertions to reflect trust policy change. --- .../src/sagemaker/core/helper/iam_policies.py | 5 +- .../train/evaluate/llm_as_judge_evaluator.py | 1 - .../train/evaluate/test_base_evaluator.py | 4 +- .../evaluate/test_bedrock_role_validation.py | 56 +++++++------------ 4 files changed, 24 insertions(+), 42 deletions(-) diff --git a/sagemaker-core/src/sagemaker/core/helper/iam_policies.py b/sagemaker-core/src/sagemaker/core/helper/iam_policies.py index 4b868bc9f3..76d7fdc26b 100644 --- a/sagemaker-core/src/sagemaker/core/helper/iam_policies.py +++ b/sagemaker-core/src/sagemaker/core/helper/iam_policies.py @@ -820,10 +820,7 @@ { "Effect": "Allow", "Principal": { - "Service": [ - "sagemaker.amazonaws.com", - "bedrock.amazonaws.com", - ] + "Service": "sagemaker.amazonaws.com" }, "Action": "sts:AssumeRole", "Condition": { diff --git a/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py b/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py index edc39ff45f..a157c05701 100644 --- a/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py +++ b/sagemaker-train/src/sagemaker/train/evaluate/llm_as_judge_evaluator.py @@ -752,7 +752,6 @@ def _get_llmaj_template_additions(self, eval_name: str) -> dict: 'evaluate_base_model': self.evaluate_base_model, } - @_telemetry_emitter( feature=Feature.MODEL_CUSTOMIZATION, func_name="LLMAsJudgeEvaluator.evaluate", diff --git a/sagemaker-train/tests/unit/train/evaluate/test_base_evaluator.py b/sagemaker-train/tests/unit/train/evaluate/test_base_evaluator.py index 9b785561c2..70307c7522 100644 --- a/sagemaker-train/tests/unit/train/evaluate/test_base_evaluator.py +++ b/sagemaker-train/tests/unit/train/evaluate/test_base_evaluator.py @@ -734,7 +734,7 @@ def test_get_aws_execution_context(self, mock_resolve, mock_role, mock_session, assert context['account_id'] == '123456789012' mock_role.assert_called_once_with( provided_role=None, - role_type="training", + role_type="evaluation", sagemaker_session=mock_session, ) @@ -761,7 +761,7 @@ def test_get_aws_execution_context_with_explicit_role(self, mock_resolve, mock_r assert context['role_arn'] == explicit_role mock_role.assert_called_once_with( provided_role=explicit_role, - role_type="training", + role_type="evaluation", sagemaker_session=mock_session, ) diff --git a/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py b/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py index 1a698bbb34..39e560ab7a 100644 --- a/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py +++ b/sagemaker-train/tests/unit/train/evaluate/test_bedrock_role_validation.py @@ -22,16 +22,21 @@ def test_evaluation_role_type_exists(self): from sagemaker.core.helper.iam_policies import IAM_POLICY_CONFIG assert "evaluation" in IAM_POLICY_CONFIG - def test_evaluation_trust_includes_bedrock(self): - """Evaluation role type should require bedrock.amazonaws.com trust.""" - expected = _expected_trust_services("evaluation") - assert "bedrock.amazonaws.com" in expected - def test_evaluation_trust_includes_sagemaker(self): """Evaluation role type should require sagemaker.amazonaws.com trust.""" expected = _expected_trust_services("evaluation") assert "sagemaker.amazonaws.com" in expected + def test_evaluation_trust_does_not_require_bedrock(self): + """Evaluation role type should NOT require bedrock.amazonaws.com trust. + + The serverless evaluation backend runs as the SageMaker execution role + and calls Bedrock APIs using the role's own credentials. Bedrock does + not need to assume the role, so trust is not required. + """ + expected = _expected_trust_services("evaluation") + assert "bedrock.amazonaws.com" not in expected + def test_evaluation_smoke_actions_include_bedrock(self): """Smoke test actions should include Bedrock evaluation actions.""" actions = _get_smoke_test_actions("evaluation") @@ -84,8 +89,8 @@ def test_resolve_passes_when_all_allowed(self): assert verdict is True assert denied == [] - def test_trust_check_fails_without_bedrock(self): - """Should return False when trust policy lacks bedrock.amazonaws.com.""" + def test_trust_check_passes_with_sagemaker(self): + """Should pass when trust policy includes sagemaker.amazonaws.com.""" mock_iam = MagicMock() mock_iam.get_role.return_value = { "Role": { @@ -100,30 +105,11 @@ def test_trust_check_fails_without_bedrock(self): } } - result = _role_trusts_service(mock_iam, "arn:aws:iam::123456789012:role/MyRole", "evaluation") - assert result is False - - def test_trust_check_passes_with_bedrock(self): - """Should return True when trust policy includes both services.""" - mock_iam = MagicMock() - mock_iam.get_role.return_value = { - "Role": { - "AssumeRolePolicyDocument": { - "Version": "2012-10-17", - "Statement": [{ - "Effect": "Allow", - "Principal": {"Service": ["sagemaker.amazonaws.com", "bedrock.amazonaws.com"]}, - "Action": "sts:AssumeRole", - }], - } - } - } - result = _role_trusts_service(mock_iam, "arn:aws:iam::123456789012:role/MyRole", "evaluation") assert result is True - def test_resolve_and_validate_raises_on_trust_failure(self): - """Full resolve_and_validate_role should raise when trust fails.""" + def test_resolve_and_validate_passes_with_sagemaker_trust(self): + """Full resolve_and_validate_role should pass with sagemaker trust only.""" with patch("sagemaker.core.helper.iam_role_resolver._get_boto_session") as mock_session: mock_boto = MagicMock() mock_session.return_value = mock_boto @@ -131,7 +117,7 @@ def test_resolve_and_validate_raises_on_trust_failure(self): mock_iam = MagicMock() mock_boto.client.return_value = mock_iam - # Role exists + # Role exists with sagemaker trust only (no bedrock needed) mock_iam.get_role.return_value = { "Role": { "Arn": "arn:aws:iam::123456789012:role/MyRole", @@ -146,7 +132,7 @@ def test_resolve_and_validate_raises_on_trust_failure(self): } } - # Permissions all allowed + # All permissions allowed actions = _get_smoke_test_actions("evaluation") paginator = MagicMock() paginator.paginate.return_value = [ @@ -154,11 +140,11 @@ def test_resolve_and_validate_raises_on_trust_failure(self): ] mock_iam.get_paginator.return_value = paginator - with pytest.raises(RoleValidationError, match="bedrock.amazonaws.com"): - resolve_and_validate_role( - provided_role="arn:aws:iam::123456789012:role/MyRole", - role_type="evaluation", - ) + result = resolve_and_validate_role( + provided_role="arn:aws:iam::123456789012:role/MyRole", + role_type="evaluation", + ) + assert result == "arn:aws:iam::123456789012:role/MyRole" def test_resolve_and_validate_passes_with_correct_role(self): """Full resolve_and_validate_role should pass with correct permissions and trust."""