Skip to content

Commit f6e41c0

Browse files
committed
tests
1 parent 1ea2927 commit f6e41c0

2 files changed

Lines changed: 173 additions & 0 deletions

File tree

Lines changed: 123 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,123 @@
1+
import os
2+
import sys
3+
from unittest.mock import MagicMock, patch
4+
5+
import pytest
6+
7+
8+
@pytest.fixture(autouse=True)
9+
def mock_sagemaker_modules():
10+
"""Mock sagemaker modules for these tests only."""
11+
# Set environment variable
12+
os.environ["SAGEMAKER_ROLE_ARN"] = "arn:aws:iam::123456789012:role/test-role"
13+
14+
# Store original modules
15+
original_modules = {}
16+
sagemaker_module_names = [
17+
"sagemaker",
18+
"sagemaker.clarify",
19+
"sagemaker.model_monitor",
20+
"sagemaker.model_monitor.dataset_format",
21+
]
22+
23+
for module_name in sagemaker_module_names:
24+
if module_name in sys.modules:
25+
original_modules[module_name] = sys.modules[module_name]
26+
27+
# Mock the sagemaker modules
28+
sys.modules["sagemaker"] = MagicMock()
29+
sys.modules["sagemaker.clarify"] = MagicMock()
30+
sys.modules["sagemaker.model_monitor"] = MagicMock()
31+
sys.modules["sagemaker.model_monitor.dataset_format"] = MagicMock()
32+
33+
# Add lambda path
34+
lambda_path = "sagemaker_model_monitoring/lambda"
35+
if lambda_path not in sys.path:
36+
sys.path.insert(0, lambda_path)
37+
38+
# Import after mocking
39+
from baselining_handler import check_baselining_job, lambda_handler
40+
41+
yield check_baselining_job, lambda_handler
42+
43+
# Restore original modules
44+
for module_name in sagemaker_module_names:
45+
if module_name in original_modules:
46+
sys.modules[module_name] = original_modules[module_name]
47+
else:
48+
sys.modules.pop(module_name, None)
49+
50+
# Remove lambda path
51+
if lambda_path in sys.path:
52+
sys.path.remove(lambda_path)
53+
54+
# Clean up imported handler module
55+
if "baselining_handler" in sys.modules:
56+
del sys.modules["baselining_handler"]
57+
58+
59+
@pytest.fixture
60+
def mock_env():
61+
"""Environment variables are already set in mock_sagemaker_modules."""
62+
pass
63+
64+
65+
@pytest.fixture
66+
def mock_sagemaker_client():
67+
"""Mock boto3 SageMaker client."""
68+
with patch("baselining_handler.sagemaker") as mock_client:
69+
yield mock_client
70+
71+
72+
def test_lambda_handler_check_action(mock_sagemaker_modules, mock_env, mock_sagemaker_client):
73+
"""Test lambda_handler with check action."""
74+
check_baselining_job, lambda_handler = mock_sagemaker_modules
75+
event = {"action": "check", "job_name": "test-job"}
76+
77+
mock_sagemaker_client.describe_processing_job.return_value = {"ProcessingJobStatus": "Completed"}
78+
79+
result = lambda_handler(event, None)
80+
81+
assert result["statusCode"] == 200
82+
assert result["status"] == "COMPLETED"
83+
assert result["job_name"] == "test-job"
84+
85+
86+
def test_check_baselining_job_completed(mock_sagemaker_modules, mock_env, mock_sagemaker_client):
87+
"""Test check_baselining_job with completed job."""
88+
check_baselining_job, lambda_handler = mock_sagemaker_modules
89+
event = {"job_name": "test-job"}
90+
91+
mock_sagemaker_client.describe_processing_job.return_value = {"ProcessingJobStatus": "Completed"}
92+
93+
result = check_baselining_job(event)
94+
95+
assert result["statusCode"] == 200
96+
assert result["status"] == "COMPLETED"
97+
assert result["job_name"] == "test-job"
98+
99+
100+
def test_check_baselining_job_in_progress(mock_sagemaker_modules, mock_env, mock_sagemaker_client):
101+
"""Test check_baselining_job with in-progress job."""
102+
check_baselining_job, lambda_handler = mock_sagemaker_modules
103+
event = {"job_name": "test-job"}
104+
105+
mock_sagemaker_client.describe_processing_job.return_value = {"ProcessingJobStatus": "InProgress"}
106+
107+
result = check_baselining_job(event)
108+
109+
assert result["statusCode"] == 200
110+
assert result["status"] == "IN_PROGRESS"
111+
assert result["job_name"] == "test-job"
112+
113+
114+
def test_check_baselining_job_missing_job_name(mock_sagemaker_modules, mock_env):
115+
"""Test check_baselining_job with missing job_name."""
116+
check_baselining_job, lambda_handler = mock_sagemaker_modules
117+
event = {}
118+
119+
result = check_baselining_job(event)
120+
121+
assert result["statusCode"] == 400
122+
assert "error" in result
123+
assert "Missing required parameter" in result["error"]

modules/sagemaker/sagemaker-model-monitoring/tests/test_stack.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,8 @@ def stack_model_package_input(
2222
enable_model_quality_monitor: bool = False,
2323
enable_model_bias_monitor: bool = False,
2424
enable_model_explainability_monitor: bool = False,
25+
baseline_training_data_s3_uri: str = None,
26+
baseline_output_data_s3_uri: str = None,
2527
) -> cdk.Stack:
2628
from sagemaker_model_monitoring import settings, stack
2729

@@ -51,6 +53,8 @@ def stack_model_package_input(
5153
enable_model_quality_monitor=enable_model_quality_monitor,
5254
enable_model_bias_monitor=enable_model_bias_monitor,
5355
enable_model_explainability_monitor=enable_model_explainability_monitor,
56+
baseline_training_data_s3_uri=baseline_training_data_s3_uri,
57+
baseline_output_data_s3_uri=baseline_output_data_s3_uri,
5458
)
5559

5660
return stack.SageMakerModelMonitoringStack(
@@ -88,6 +92,36 @@ def test_synthesize_stack_model_explainability(stack_defaults: None) -> None:
8892
template.resource_count_is("AWS::SageMaker::ModelExplainabilityJobDefinition", 1)
8993

9094

95+
def test_baseline_generation_resources(stack_defaults: None) -> None:
96+
"""Test that baseline generation resources are created when URIs are provided."""
97+
stack = stack_model_package_input(
98+
enable_data_quality_monitor=True,
99+
baseline_training_data_s3_uri="s3://test-bucket/training-data.csv",
100+
baseline_output_data_s3_uri="s3://test-bucket/baselines/",
101+
)
102+
template = Template.from_stack(stack)
103+
104+
# Check for Lambda function
105+
template.resource_count_is("AWS::Lambda::Function", 1)
106+
107+
# Check for Step Functions state machine
108+
template.resource_count_is("AWS::StepFunctions::StateMachine", 1)
109+
110+
# Check for EventBridge rule
111+
template.resource_count_is("AWS::Events::Rule", 1)
112+
113+
114+
def test_no_baseline_generation_without_uris(stack_defaults: None) -> None:
115+
"""Test that baseline generation resources are NOT created without URIs."""
116+
stack = stack_model_package_input(enable_data_quality_monitor=True)
117+
template = Template.from_stack(stack)
118+
119+
# Should not have Lambda, Step Functions, or EventBridge resources
120+
template.resource_count_is("AWS::Lambda::Function", 0)
121+
template.resource_count_is("AWS::StepFunctions::StateMachine", 0)
122+
template.resource_count_is("AWS::Events::Rule", 0)
123+
124+
91125
def test_no_cdk_nag_errors(stack_defaults: None) -> None:
92126
stack = stack_model_package_input(
93127
enable_data_quality_monitor=True,
@@ -102,3 +136,19 @@ def test_no_cdk_nag_errors(stack_defaults: None) -> None:
102136
Match.string_like_regexp(r"AwsSolutions-.*"),
103137
)
104138
assert not nag_errors, f"Found {len(nag_errors)} CDK nag errors"
139+
140+
141+
def test_no_cdk_nag_errors_with_baseline(stack_defaults: None) -> None:
142+
"""Test CDK nag with baseline generation enabled."""
143+
stack = stack_model_package_input(
144+
enable_data_quality_monitor=True,
145+
baseline_training_data_s3_uri="s3://test-bucket/training-data.csv",
146+
baseline_output_data_s3_uri="s3://test-bucket/baselines/",
147+
)
148+
cdk.Aspects.of(stack).add(cdk_nag.AwsSolutionsChecks())
149+
150+
nag_errors = Annotations.from_stack(stack).find_error(
151+
"*",
152+
Match.string_like_regexp(r"AwsSolutions-.*"),
153+
)
154+
assert not nag_errors, f"Found {len(nag_errors)} CDK nag errors"

0 commit comments

Comments
 (0)