Files
mlflow--mlflow/tests/sagemaker/mock/test_sagemaker_service_mock.py
2026-07-13 13:22:34 +08:00

682 lines
24 KiB
Python

import boto3
import pytest
from tests.helper_functions import set_boto_credentials # noqa: F401
from tests.sagemaker.mock import mock_sagemaker
@pytest.fixture
def sagemaker_client():
return boto3.client("sagemaker", region_name="us-west-2")
def create_sagemaker_model(sagemaker_client, model_name):
return sagemaker_client.create_model(
ExecutionRoleArn="arn:aws:iam::012345678910:role/sample-role",
ModelName=model_name,
PrimaryContainer={
"Image": "012345678910.dkr.ecr.us-west-2.amazonaws.com/sample-container",
},
)
def create_endpoint_config(sagemaker_client, endpoint_config_name, model_name):
return sagemaker_client.create_endpoint_config(
EndpointConfigName=endpoint_config_name,
ProductionVariants=[
{
"VariantName": "sample-variant",
"ModelName": model_name,
"InitialInstanceCount": 1,
"InstanceType": "ml.m4.xlarge",
"InitialVariantWeight": 1.0,
},
],
)
@mock_sagemaker
def test_created_model_is_listed_by_list_models_function(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
models_response = sagemaker_client.list_models()
assert "Models" in models_response
models = models_response["Models"]
assert all("ModelName" in model for model in models)
assert model_name in [model["ModelName"] for model in models]
@mock_sagemaker
def test_create_model_returns_arn_containing_model_name(sagemaker_client):
model_name = "sample-model"
model_create_response = create_sagemaker_model(
sagemaker_client=sagemaker_client, model_name=model_name
)
assert "ModelArn" in model_create_response
assert model_name in model_create_response["ModelArn"]
@mock_sagemaker
def test_creating_model_with_name_already_in_use_raises_exception(sagemaker_client):
model_name = "sample-model-name"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
with pytest.raises(ValueError, match="Attempted to create a model"):
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
@mock_sagemaker
def test_all_models_are_listed_after_creating_many_models(sagemaker_client):
model_names = []
for i in range(100):
model_name = f"sample-model-{i}"
model_names.append(model_name)
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
listed_models = sagemaker_client.list_models()["Models"]
listed_model_names = [model["ModelName"] for model in listed_models]
for model_name in model_names:
assert model_name in listed_model_names
@mock_sagemaker
def test_describe_model_response_contains_expected_attributes(sagemaker_client):
model_name = "sample-model"
execution_role_arn = "arn:aws:iam::012345678910:role/sample-role"
primary_container = {
"Image": "012345678910.dkr.ecr.us-west-2.amazonaws.com/sample-container",
}
sagemaker_client.create_model(
ModelName=model_name,
ExecutionRoleArn=execution_role_arn,
PrimaryContainer=primary_container,
)
describe_model_response = sagemaker_client.describe_model(ModelName=model_name)
assert "CreationTime" in describe_model_response
assert "ModelArn" in describe_model_response
assert "ExecutionRoleArn" in describe_model_response
assert describe_model_response["ExecutionRoleArn"] == execution_role_arn
assert "ModelName" in describe_model_response
assert describe_model_response["ModelName"] == model_name
assert "PrimaryContainer" in describe_model_response
assert describe_model_response["PrimaryContainer"] == primary_container
@mock_sagemaker
def test_describe_model_throws_exception_for_nonexistent_model(sagemaker_client):
with pytest.raises(ValueError, match="Attempted to describe a model"):
sagemaker_client.describe_model(ModelName="nonexistent-model")
@mock_sagemaker
def test_model_is_no_longer_listed_after_deletion(sagemaker_client):
model_name = "sample-model-name"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
sagemaker_client.delete_model(ModelName=model_name)
listed_models = sagemaker_client.list_models()["Models"]
listed_model_names = [model["ModelName"] for model in listed_models]
assert model_name not in listed_model_names
@mock_sagemaker
def test_created_endpoint_config_is_listed_by_list_endpoints_function(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
endpoint_configs_response = sagemaker_client.list_endpoint_configs()
assert "EndpointConfigs" in endpoint_configs_response
endpoint_configs = endpoint_configs_response["EndpointConfigs"]
assert all("EndpointConfigName" in endpoint_config for endpoint_config in endpoint_configs)
assert endpoint_config_name in [
endpoint_config["EndpointConfigName"] for endpoint_config in endpoint_configs
]
@mock_sagemaker
def test_create_endpoint_config_returns_arn_containing_config_name(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_config_response = create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
assert "EndpointConfigArn" in create_config_response
assert endpoint_config_name in create_config_response["EndpointConfigArn"]
@mock_sagemaker
def test_creating_endpoint_config_with_name_already_in_use_raises_exception(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
with pytest.raises(ValueError, match="Attempted to create an endpoint configuration"):
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
@mock_sagemaker
def test_all_endpoint_configs_are_listed_after_creating_many_configs(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_names = []
for i in range(100):
endpoint_config_name = f"sample-config-{i}"
endpoint_config_names.append(endpoint_config_name)
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
listed_endpoint_configs = sagemaker_client.list_endpoint_configs()["EndpointConfigs"]
listed_endpoint_config_names = [
endpoint_config["EndpointConfigName"] for endpoint_config in listed_endpoint_configs
]
for endpoint_config_name in endpoint_config_names:
assert endpoint_config_name in listed_endpoint_config_names
@mock_sagemaker
def test_describe_endpoint_config_response_contains_expected_attributes(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
production_variants = [
{
"VariantName": "sample-variant",
"ModelName": model_name,
"InitialInstanceCount": 1,
"InstanceType": "ml.m4.xlarge",
"InitialVariantWeight": 1.0,
},
]
async_inference_config = {
"ClientConfig": {"MaxConcurrentInvocationsPerInstance": 4},
"OutputConfig": {"S3OutputPath": "s3://bucket_name/", "NotificationConfig": {}},
}
sagemaker_client.create_endpoint_config(
EndpointConfigName=endpoint_config_name,
ProductionVariants=production_variants,
AsyncInferenceConfig=async_inference_config,
)
describe_endpoint_config_response = sagemaker_client.describe_endpoint_config(
EndpointConfigName=endpoint_config_name
)
assert "CreationTime" in describe_endpoint_config_response
assert "EndpointConfigArn" in describe_endpoint_config_response
assert "EndpointConfigName" in describe_endpoint_config_response
assert describe_endpoint_config_response["EndpointConfigName"] == endpoint_config_name
assert "ProductionVariants" in describe_endpoint_config_response
assert describe_endpoint_config_response["ProductionVariants"] == production_variants
assert "AsyncInferenceConfig" in describe_endpoint_config_response
assert describe_endpoint_config_response["AsyncInferenceConfig"] == async_inference_config
@mock_sagemaker
def test_describe_endpoint_config_throws_exception_for_nonexistent_config(sagemaker_client):
with pytest.raises(ValueError, match="Attempted to describe an endpoint config"):
sagemaker_client.describe_endpoint_config(EndpointConfigName="nonexistent-config")
@mock_sagemaker
def test_endpoint_config_is_no_longer_listed_after_deletion(sagemaker_client):
model_name = "sample-model-name"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
sagemaker_client.delete_endpoint_config(EndpointConfigName=endpoint_config_name)
listed_endpoint_configs = sagemaker_client.list_endpoint_configs()["EndpointConfigs"]
listed_endpoint_config_names = [
endpoint_config["EndpointConfigName"] for endpoint_config in listed_endpoint_configs
]
assert endpoint_config_name not in listed_endpoint_config_names
@mock_sagemaker
def test_created_endpoint_is_listed_by_list_endpoints_function(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
endpoint_name = "sample-endpoint"
sagemaker_client.create_endpoint(
EndpointConfigName=endpoint_config_name,
EndpointName=endpoint_name,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
endpoints_response = sagemaker_client.list_endpoints()
assert "Endpoints" in endpoints_response
endpoints = endpoints_response["Endpoints"]
assert all("EndpointName" in endpoint for endpoint in endpoints)
assert endpoint_name in [endpoint["EndpointName"] for endpoint in endpoints]
@mock_sagemaker
def test_create_endpoint_returns_arn_containing_endpoint_name(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
endpoint_name = "sample-endpoint"
create_endpoint_response = sagemaker_client.create_endpoint(
EndpointConfigName=endpoint_config_name,
EndpointName=endpoint_name,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
assert "EndpointArn" in create_endpoint_response
assert endpoint_name in create_endpoint_response["EndpointArn"]
@mock_sagemaker
def test_creating_endpoint_with_name_already_in_use_raises_exception(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
endpoint_name = "sample-endpoint"
sagemaker_client.create_endpoint(
EndpointConfigName=endpoint_config_name,
EndpointName=endpoint_name,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
with pytest.raises(ValueError, match="Attempted to create an endpoint"):
sagemaker_client.create_endpoint(
EndpointConfigName=endpoint_config_name,
EndpointName=endpoint_name,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
@mock_sagemaker
def test_all_endpoint_are_listed_after_creating_many_endpoints(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
endpoint_names = []
for i in range(100):
endpoint_name = f"sample-endpoint-{i}"
endpoint_names.append(endpoint_name)
sagemaker_client.create_endpoint(
EndpointConfigName=endpoint_config_name,
EndpointName=endpoint_name,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
listed_endpoints = sagemaker_client.list_endpoints()["Endpoints"]
listed_endpoint_names = [endpoint["EndpointName"] for endpoint in listed_endpoints]
for endpoint_name in endpoint_names:
assert endpoint_name in listed_endpoint_names
@mock_sagemaker
def test_describe_endpoint_response_contains_expected_attributes(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
production_variants = [
{
"VariantName": "sample-variant",
"ModelName": model_name,
"InitialInstanceCount": 1,
"InstanceType": "ml.m4.xlarge",
"InitialVariantWeight": 1.0,
},
]
sagemaker_client.create_endpoint_config(
EndpointConfigName=endpoint_config_name,
ProductionVariants=production_variants,
)
endpoint_name = "sample-endpoint"
sagemaker_client.create_endpoint(
EndpointName=endpoint_name,
EndpointConfigName=endpoint_config_name,
)
describe_endpoint_response = sagemaker_client.describe_endpoint(EndpointName=endpoint_name)
assert "CreationTime" in describe_endpoint_response
assert "LastModifiedTime" in describe_endpoint_response
assert "EndpointArn" in describe_endpoint_response
assert "EndpointStatus" in describe_endpoint_response
assert "ProductionVariants" in describe_endpoint_response
@mock_sagemaker
def test_describe_endpoint_throws_exception_for_nonexistent_endpoint(sagemaker_client):
with pytest.raises(ValueError, match="Attempted to describe an endpoint"):
sagemaker_client.describe_endpoint(EndpointName="nonexistent-endpoint")
@mock_sagemaker
def test_endpoint_is_no_longer_listed_after_deletion(sagemaker_client):
model_name = "sample-model-name"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
endpoint_name = "sample-endpoint"
sagemaker_client.create_endpoint(
EndpointConfigName=endpoint_config_name,
EndpointName=endpoint_name,
)
sagemaker_client.delete_endpoint(EndpointName=endpoint_name)
listed_endpoints = sagemaker_client.list_endpoints()["Endpoints"]
listed_endpoint_names = [endpoint["EndpointName"] for endpoint in listed_endpoints]
assert endpoint_name not in listed_endpoint_names
@mock_sagemaker
def test_update_endpoint_modifies_config_correctly(sagemaker_client):
model_name = "sample-model-name"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
first_endpoint_config_name = "sample-config-1"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=first_endpoint_config_name,
model_name=model_name,
)
second_endpoint_config_name = "sample-config-2"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=second_endpoint_config_name,
model_name=model_name,
)
endpoint_name = "sample-endpoint"
sagemaker_client.create_endpoint(
EndpointConfigName=first_endpoint_config_name,
EndpointName=endpoint_name,
)
first_describe_endpoint_response = sagemaker_client.describe_endpoint(
EndpointName=endpoint_name
)
assert first_describe_endpoint_response["EndpointConfigName"] == first_endpoint_config_name
sagemaker_client.update_endpoint(
EndpointName=endpoint_name, EndpointConfigName=second_endpoint_config_name
)
second_describe_endpoint_response = sagemaker_client.describe_endpoint(
EndpointName=endpoint_name
)
assert second_describe_endpoint_response["EndpointConfigName"] == second_endpoint_config_name
@mock_sagemaker
def test_update_endpoint_with_nonexistent_config_throws_exception(sagemaker_client):
model_name = "sample-model-name"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
endpoint_config_name = "sample-config"
create_endpoint_config(
sagemaker_client=sagemaker_client,
endpoint_config_name=endpoint_config_name,
model_name=model_name,
)
endpoint_name = "sample-endpoint"
sagemaker_client.create_endpoint(
EndpointConfigName=endpoint_config_name,
EndpointName=endpoint_name,
)
with pytest.raises(ValueError, match="Attempted to update an endpoint"):
sagemaker_client.update_endpoint(
EndpointName=endpoint_name, EndpointConfigName="nonexistent-config"
)
@mock_sagemaker
def test_created_transform_job_is_listed_by_list_transform_jobs_function(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
transform_input = {
"DataSource": {"S3DataSource": {"S3DataType": "Some Data Type", "S3Uri": "Some Input Uri"}}
}
transform_output = {"S3OutputPath": "Some Output Path"}
transform_resources = {"InstanceType": "Some Instance Type", "InstanceCount": 1}
job_name = "sample-job"
sagemaker_client.create_transform_job(
TransformJobName=job_name,
ModelName=model_name,
TransformInput=transform_input,
TransformOutput=transform_output,
TransformResources=transform_resources,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
transform_jobs_response = sagemaker_client.list_transform_jobs()
assert "TransformJobSummaries" in transform_jobs_response
transform_jobs = transform_jobs_response["TransformJobSummaries"]
assert all("TransformJobName" in transform_job for transform_job in transform_jobs)
assert job_name in [transform_job["TransformJobName"] for transform_job in transform_jobs]
@mock_sagemaker
def test_create_transform_job_returns_arn_containing_transform_job_name(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
transform_input = {
"DataSource": {"S3DataSource": {"S3DataType": "Some Data Type", "S3Uri": "Some Input Uri"}}
}
transform_output = {"S3OutputPath": "Some Output Path"}
transform_resources = {"InstanceType": "Some Instance Type", "InstanceCount": 1}
job_name = "sample-job"
create_transform_job_response = sagemaker_client.create_transform_job(
TransformJobName=job_name,
ModelName=model_name,
TransformInput=transform_input,
TransformOutput=transform_output,
TransformResources=transform_resources,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
assert "TransformJobArn" in create_transform_job_response
assert job_name in create_transform_job_response["TransformJobArn"]
@mock_sagemaker
def test_creating_transform_job_with_name_already_in_use_raises_exception(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
transform_input = {
"DataSource": {"S3DataSource": {"S3DataType": "Some Data Type", "S3Uri": "Some Input Uri"}}
}
transform_output = {"S3OutputPath": "Some Output Path"}
transform_resources = {"InstanceType": "Some Instance Type", "InstanceCount": 1}
job_name = "sample-job"
sagemaker_client.create_transform_job(
TransformJobName=job_name,
ModelName=model_name,
TransformInput=transform_input,
TransformOutput=transform_output,
TransformResources=transform_resources,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
with pytest.raises(ValueError, match="Attempted to create a transform job"):
sagemaker_client.create_transform_job(
TransformJobName=job_name,
ModelName=model_name,
TransformInput=transform_input,
TransformOutput=transform_output,
TransformResources=transform_resources,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
@mock_sagemaker
def test_all_transform_jobs_are_listed_after_creating_many_transform_jobs(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
transform_input = {
"DataSource": {"S3DataSource": {"S3DataType": "Some Data Type", "S3Uri": "Some Input Uri"}}
}
transform_output = {"S3OutputPath": "Some Output Path"}
transform_resources = {"InstanceType": "Some Instance Type", "InstanceCount": 1}
job_names = []
for i in range(100):
job_name = f"sample-job-{i}"
job_names.append(job_name)
sagemaker_client.create_transform_job(
TransformJobName=job_name,
ModelName=model_name,
TransformInput=transform_input,
TransformOutput=transform_output,
TransformResources=transform_resources,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
listed_transform_jobs = sagemaker_client.list_transform_jobs()["TransformJobSummaries"]
listed_transform_job_names = [
transform_job["TransformJobName"] for transform_job in listed_transform_jobs
]
for job_name in job_names:
assert job_name in listed_transform_job_names
@mock_sagemaker
def test_describe_transform_job_response_contains_expected_attributes(sagemaker_client):
model_name = "sample-model"
create_sagemaker_model(sagemaker_client=sagemaker_client, model_name=model_name)
transform_input = {
"DataSource": {"S3DataSource": {"S3DataType": "Some Data Type", "S3Uri": "Some Input Uri"}}
}
transform_output = {"S3OutputPath": "Some Output Path"}
transform_resources = {"InstanceType": "Some Instance Type", "InstanceCount": 1}
job_name = "sample-job"
sagemaker_client.create_transform_job(
TransformJobName=job_name,
ModelName=model_name,
TransformInput=transform_input,
TransformOutput=transform_output,
TransformResources=transform_resources,
Tags=[{"Key": "Some Key", "Value": "Some Value"}],
)
describe_transform_job_response = sagemaker_client.describe_transform_job(
TransformJobName=job_name
)
assert "TransformJobName" in describe_transform_job_response
assert "CreationTime" in describe_transform_job_response
assert "TransformJobArn" in describe_transform_job_response
assert "TransformJobStatus" in describe_transform_job_response
assert "ModelName" in describe_transform_job_response
@mock_sagemaker
def test_describe_transform_job_throws_exception_for_nonexistent_transform_job(sagemaker_client):
with pytest.raises(ValueError, match="Attempted to describe a transform job"):
sagemaker_client.describe_transform_job(TransformJobName="nonexistent-job")