From f4905494264d749fa2b3a8dace925dbaa3f2add6 Mon Sep 17 00:00:00 2001 From: Rishabh Devnani Date: Sat, 8 Aug 2026 00:03:23 +0000 Subject: [PATCH] feat(pipeline): Add inference and lineage step types Add 4 pipeline step classes: - EndpointConfigStep, EndpointStep (SageMaker inference deployment) - InferenceComponentStep (multi-model endpoint support) - LineageStep (ML governance tracking) Design: each step accepts an 'arguments: Dict[str, Any]' forwarded to the pipeline service. Top-level argument keys are validated client-side against the corresponding public AWS API input shape (botocore service model) at construction and at serialization; fields the service is known to reject fail fast with actionable errors (EndpointConfig: DataCaptureConfig, ExplainerConfig; Endpoint: DeploymentConfig). Values are not validated -- they may be pipeline variables resolved at compile time. Full schema validation remains server-side. If the installed botocore does not know an operation, shape validation is skipped and the service remains the authority. Retryability: only EndpointConfigStep is retryable. Cacheability: EndpointConfigStep and EndpointStep are structurally cacheable via cache_config. Includes 23 unit tests and a LineageStep end-to-end integration test. --- X-AI-Prompt: Add the inference and lineage pipeline step types to the Python SDK with client-side argument validation X-AI-Tool: kiro-cli --- .../src/sagemaker/mlops/workflow/__init__.py | 8 + .../mlops/workflow/_argument_validation.py | 124 +++++++ .../sagemaker/mlops/workflow/endpoint_step.py | 223 ++++++++++++ .../workflow/inference_component_step.py | 107 ++++++ .../sagemaker/mlops/workflow/lineage_step.py | 117 ++++++ .../src/sagemaker/mlops/workflow/steps.py | 4 + .../tests/integ/workflow/test_lineage_step.py | 145 ++++++++ .../workflow/test_inference_lineage_steps.py | 332 ++++++++++++++++++ 8 files changed, 1060 insertions(+) create mode 100644 sagemaker-mlops/src/sagemaker/mlops/workflow/_argument_validation.py create mode 100644 sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py create mode 100644 sagemaker-mlops/src/sagemaker/mlops/workflow/inference_component_step.py create mode 100644 sagemaker-mlops/src/sagemaker/mlops/workflow/lineage_step.py create mode 100644 sagemaker-mlops/tests/integ/workflow/test_lineage_step.py create mode 100644 sagemaker-mlops/tests/unit/workflow/test_inference_lineage_steps.py diff --git a/sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py b/sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py index 129abb1c76..a5ab5ba851 100644 --- a/sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py +++ b/sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py @@ -14,6 +14,7 @@ functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve for orchestration purposes. """ + from __future__ import absolute_import __version__ = "0.1.0" @@ -46,8 +47,11 @@ from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep from sagemaker.mlops.workflow.condition_step import ConditionStep from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig +from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep from sagemaker.mlops.workflow.fail_step import FailStep +from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput +from sagemaker.mlops.workflow.lineage_step import LineageStep from sagemaker.mlops.workflow.model_step import ModelStep from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep @@ -98,9 +102,13 @@ "ConditionStep", "EMRStep", "EMRStepConfig", + "EndpointConfigStep", + "EndpointStep", "FailStep", + "InferenceComponentStep", "LambdaStep", "LambdaOutput", + "LineageStep", "ModelStep", "MonitorBatchTransformStep", "NotebookJobStep", diff --git a/sagemaker-mlops/src/sagemaker/mlops/workflow/_argument_validation.py b/sagemaker-mlops/src/sagemaker/mlops/workflow/_argument_validation.py new file mode 100644 index 0000000000..24072c4d4a --- /dev/null +++ b/sagemaker-mlops/src/sagemaker/mlops/workflow/_argument_validation.py @@ -0,0 +1,124 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Client-side validation for pipeline step ``arguments`` blocks. + +Validates the **top-level keys** of a step's ``arguments`` dict against +the corresponding public AWS API input shape from botocore, and rejects +fields that SageMaker Pipelines is known not to support. This fails fast +at step construction with a clear error, instead of a server-side parse +failure at ``CreatePipeline`` time. + +Values are intentionally not validated: they may be pipeline variables +(parameter references, step property references, ``Join``/``JsonGet`` +expressions) that only resolve at pipeline compile or execution time. + +If the installed botocore release does not know the target operation +(for example, a very old botocore release), shape +validation is skipped and the service remains the authority. +""" + +from __future__ import absolute_import + +import logging +from typing import Any, Dict, FrozenSet, Optional, Sequence, Tuple + +import botocore.session +from botocore.exceptions import UnknownServiceError +from botocore.model import OperationNotFoundError + +logger = logging.getLogger(__name__) + +# Cache of (service, operation) -> allowed top-level keys. +# ``None`` means botocore does not know the operation; skip shape checks. +_SHAPE_CACHE: Dict[Tuple[str, str], Optional[FrozenSet[str]]] = {} + + +def _allowed_top_level_keys(service_name: str, operation_name: str) -> Optional[FrozenSet[str]]: + """Return the allowed top-level keys for an operation input shape. + + Args: + service_name (str): botocore service name (e.g. ``sagemaker``). + operation_name (str): operation name (e.g. ``CreateEndpointConfig``). + + Returns: + The allowed key set, or ``None`` if the installed botocore does + not know the operation (validation should then be skipped). + """ + cache_key = (service_name, operation_name) + if cache_key not in _SHAPE_CACHE: + try: + session = botocore.session.get_session() + service_model = session.get_service_model(service_name) + operation_model = service_model.operation_model(operation_name) + members = operation_model.input_shape.members.keys() + _SHAPE_CACHE[cache_key] = frozenset(members) + except (UnknownServiceError, OperationNotFoundError): + logger.warning( + "Installed botocore does not know %s.%s; skipping " + "client-side argument shape validation for this step.", + service_name, + operation_name, + ) + _SHAPE_CACHE[cache_key] = None + return _SHAPE_CACHE[cache_key] + + +def validate_step_arguments( + step_class_name: str, + arguments: Dict[str, Any], + service_name: str, + operation_name: str, + unsupported_fields: Sequence[str] = (), +) -> None: + """Validate the top-level keys of a step ``arguments`` dict. + + Args: + step_class_name (str): Step class name, used in error messages. + arguments (Dict[str, Any]): The user-provided ``arguments`` dict. + service_name (str): botocore service name of the wrapped API. + operation_name (str): Operation whose input shape defines the + allowed top-level fields. + unsupported_fields (Sequence[str]): Fields that exist in the + public API shape but are rejected by SageMaker Pipelines. + + Raises: + ValueError: If ``arguments`` is not a non-empty dict with string + keys, contains an unsupported field, or contains a key that + is not part of the operation's input shape. + """ + if arguments is None: + raise ValueError(f"arguments is required for {step_class_name}.") + if not isinstance(arguments, dict) or not arguments: + raise ValueError(f"{step_class_name}: arguments must be a non-empty dict.") + non_string_keys = [key for key in arguments if not isinstance(key, str)] + if non_string_keys: + raise ValueError( + f"{step_class_name}: argument keys must be strings; got {non_string_keys!r}." + ) + rejected = sorted(field for field in unsupported_fields if field in arguments) + if rejected: + raise ValueError( + f"{step_class_name}: field(s) {rejected} are not supported by " + "SageMaker Pipelines and would be rejected at pipeline creation " + "time. Remove them from arguments." + ) + allowed = _allowed_top_level_keys(service_name, operation_name) + if allowed is None: + return + unknown = sorted(set(arguments) - allowed) + if unknown: + raise ValueError( + f"{step_class_name}: unknown argument field(s) {unknown}. " + f"Allowed top-level fields (from {service_name}.{operation_name}): " + f"{sorted(allowed)}." + ) diff --git a/sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py b/sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py new file mode 100644 index 0000000000..b3588f359a --- /dev/null +++ b/sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py @@ -0,0 +1,223 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Step definitions for SageMaker Endpoint deployment in Pipelines. + +Design note: the pipeline service models each step's +``Arguments`` block as an opaque structure validated against +the underlying SageMaker request model (``CreateEndpointConfigInput`` +or ``CreateEndpointInput``) minus a small exclusion set. This SDK +validates the **top-level keys** of the ``arguments`` dict against the +public ``CreateEndpointConfig``/``CreateEndpoint`` API input shape at +construction time (values are not validated -- they may be pipeline +variables) and forwards the dict to the service, which remains the +authority on full schema validation. + +Excluded fields (the pipeline service rejects the pipeline if present): + +* ``EndpointConfig``: ``DataCaptureConfig``, ``ExplainerConfig`` +* ``Endpoint``: ``DeploymentConfig`` +""" + +from __future__ import absolute_import + +from typing import Any, Dict, List, Optional, Union + +from sagemaker.core.helper.pipeline_variable import RequestType +from sagemaker.core.workflow.properties import Properties + +from sagemaker.mlops.workflow._argument_validation import validate_step_arguments +from sagemaker.mlops.workflow.retry import RetryPolicy +from sagemaker.mlops.workflow.step_collections import StepCollection +from sagemaker.mlops.workflow.steps import ( + CacheConfig, + ConfigurableRetryStep, + Step, + StepTypeEnum, +) + + +class EndpointConfigStep(ConfigurableRetryStep): + """Creates a SageMaker EndpointConfig within a pipeline. + + Wraps the SageMaker ``CreateEndpointConfig`` API. The ``arguments`` + dict is passed through to the service; it accepts any field of + ``CreateEndpointConfigInput`` **except** ``DataCaptureConfig`` and + ``ExplainerConfig``, which are rejected by the pipeline service. + + Per the pipeline service's step contract, ``EndpointConfig`` is structurally + cacheable (``cache_config``) and retryable (``retry_policies``). + """ + + def __init__( + self, + name: str, + arguments: Dict[str, Any], + display_name: Optional[str] = None, + description: Optional[str] = None, + depends_on: Optional[List[Union[str, Step, StepCollection]]] = None, + cache_config: Optional[CacheConfig] = None, + retry_policies: Optional[List[RetryPolicy]] = None, + ): + """Construct an ``EndpointConfigStep``. + + Args: + name (str): The name of the step. + arguments (Dict[str, Any]): The ``Arguments`` block for the + ``CreateEndpointConfig`` call. Required fields: + ``EndpointConfigName``, ``ProductionVariants``. Optional + fields include ``KmsKeyId``, ``AsyncInferenceConfig``, + ``ShadowProductionVariants``, ``ExecutionRoleArn``, + ``VpcConfig``, ``EnableNetworkIsolation``, + ``MetricsConfig``. Values may be pipeline variables + (parameter references, step property references) — the + pipeline compiler resolves them at definition time. + Do not include ``DataCaptureConfig`` or ``ExplainerConfig`` + (the pipeline service rejects them). + display_name (str): Optional display name. + description (str): Optional description. + depends_on (List[Union[str, Step, StepCollection]]): Optional + explicit step dependencies. + cache_config (CacheConfig): Optional cache configuration. + retry_policies (List[RetryPolicy]): Optional retry policies. + """ + super().__init__( + name=name, + step_type=StepTypeEnum.ENDPOINT_CONFIG, + display_name=display_name, + description=description, + depends_on=depends_on, + retry_policies=retry_policies, + ) + if arguments is None: + raise ValueError("arguments is required for EndpointConfigStep.") + validate_step_arguments( + "EndpointConfigStep", + arguments, + service_name="sagemaker", + operation_name="CreateEndpointConfig", + unsupported_fields=("DataCaptureConfig", "ExplainerConfig"), + ) + self._arguments = arguments + self.cache_config = cache_config + self._properties = Properties( + step_name=name, step=self, shape_name="DescribeEndpointConfigOutput" + ) + + @property + def arguments(self) -> RequestType: + """The ``Arguments`` block for the ``CreateEndpointConfig`` call.""" + validate_step_arguments( + "EndpointConfigStep", + self._arguments, + service_name="sagemaker", + operation_name="CreateEndpointConfig", + unsupported_fields=("DataCaptureConfig", "ExplainerConfig"), + ) + return self._arguments + + @property + def properties(self): + """A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``.""" + return self._properties + + def to_request(self) -> RequestType: + """Get the request structure for workflow service calls.""" + request_dict = super().to_request() + if self.cache_config: + request_dict.update(self.cache_config.config) + return request_dict + + +class EndpointStep(Step): + """Creates or updates a SageMaker Endpoint within a pipeline. + + Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API — the + pipeline chooses create-vs-update based on endpoint existence. The + ``arguments`` dict is passed through to the service; it accepts any + field of ``CreateEndpointInput`` **except** ``DeploymentConfig``, + which is rejected by the pipeline service. + + Per the pipeline service's step contract, ``Endpoint`` is structurally cacheable + but not retryable at the pipeline level. + """ + + def __init__( + self, + name: str, + arguments: Dict[str, Any], + display_name: Optional[str] = None, + description: Optional[str] = None, + depends_on: Optional[List[Union[str, Step, StepCollection]]] = None, + cache_config: Optional[CacheConfig] = None, + ): + """Construct an ``EndpointStep``. + + Args: + name (str): The name of the step. + arguments (Dict[str, Any]): The ``Arguments`` block for the + ``CreateEndpoint`` / ``UpdateEndpoint`` call. Required + fields: ``EndpointName``, ``EndpointConfigName``. Optional + fields: ``GraphConfigName``, ``DeletionCondition``. + Values may be pipeline variables. Do not include + ``DeploymentConfig`` (the pipeline service rejects it). + display_name (str): Optional display name. + description (str): Optional description. + depends_on (List[Union[str, Step, StepCollection]]): Optional + explicit step dependencies. + cache_config (CacheConfig): Optional cache configuration. + """ + super().__init__( + name=name, + display_name=display_name, + description=description, + step_type=StepTypeEnum.ENDPOINT, + depends_on=depends_on, + ) + if arguments is None: + raise ValueError("arguments is required for EndpointStep.") + validate_step_arguments( + "EndpointStep", + arguments, + service_name="sagemaker", + operation_name="CreateEndpoint", + unsupported_fields=("DeploymentConfig",), + ) + self._arguments = arguments + self.cache_config = cache_config + self._properties = Properties( + step_name=name, step=self, shape_name="DescribeEndpointOutput" + ) + + @property + def arguments(self) -> RequestType: + """The ``Arguments`` block for the ``CreateEndpoint``/``UpdateEndpoint`` call.""" + validate_step_arguments( + "EndpointStep", + self._arguments, + service_name="sagemaker", + operation_name="CreateEndpoint", + unsupported_fields=("DeploymentConfig",), + ) + return self._arguments + + @property + def properties(self): + """A ``Properties`` object shaped like ``DescribeEndpointOutput``.""" + return self._properties + + def to_request(self) -> RequestType: + """Get the request structure for workflow service calls.""" + request_dict = super().to_request() + if self.cache_config: + request_dict.update(self.cache_config.config) + return request_dict diff --git a/sagemaker-mlops/src/sagemaker/mlops/workflow/inference_component_step.py b/sagemaker-mlops/src/sagemaker/mlops/workflow/inference_component_step.py new file mode 100644 index 0000000000..935aab1078 --- /dev/null +++ b/sagemaker-mlops/src/sagemaker/mlops/workflow/inference_component_step.py @@ -0,0 +1,107 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Step definition for SageMaker InferenceComponent in Pipelines. + +Design note: the ``Arguments`` block is validated server-side against SageMaker's +``CreateInferenceComponentInput`` request model with no field +exclusions — any field the AWS API accepts, the pipeline service +accepts. +""" + +from __future__ import absolute_import + +from typing import Any, Dict, List, Optional, Union + +from sagemaker.core.helper.pipeline_variable import RequestType +from sagemaker.core.workflow.properties import Properties + +from sagemaker.mlops.workflow._argument_validation import validate_step_arguments +from sagemaker.mlops.workflow.step_collections import StepCollection +from sagemaker.mlops.workflow.steps import Step, StepTypeEnum + + +class InferenceComponentStep(Step): + """Creates or updates a SageMaker Inference Component within a pipeline. + + Wraps the SageMaker ``CreateInferenceComponent``/``UpdateInferenceComponent`` + API — the pipeline chooses create-vs-update based on component existence. + Inference components enable multi-model endpoint deployments with + independent scaling per model. + + The ``arguments`` dict is passed through to the service; it accepts + any field of ``CreateInferenceComponentInput`` (no exclusions). + + Per the pipeline service's step contract, ``InferenceComponent`` is neither + cacheable nor retryable at the pipeline level. + """ + + def __init__( + self, + name: str, + arguments: Dict[str, Any], + display_name: Optional[str] = None, + description: Optional[str] = None, + depends_on: Optional[List[Union[str, Step, StepCollection]]] = None, + ): + """Construct an ``InferenceComponentStep``. + + Args: + name (str): The name of the step. + arguments (Dict[str, Any]): The ``Arguments`` block for the + ``CreateInferenceComponent``/``UpdateInferenceComponent`` + call. Typical fields: ``InferenceComponentName``, + ``EndpointName``, ``VariantName``, ``Specification``, + ``Specifications`` (plural, for multi-spec deployments), + ``RuntimeConfig``. Values may be pipeline variables. + Note: ``ComputeResourceRequirements.NumberOfCpuCoresRequired`` + is a float — pass ``2.0`` not ``2``. + display_name (str): Optional display name. + description (str): Optional description. + depends_on (List[Union[str, Step, StepCollection]]): Optional + explicit step dependencies. + """ + super().__init__( + name=name, + display_name=display_name, + description=description, + step_type=StepTypeEnum.INFERENCE_COMPONENT, + depends_on=depends_on, + ) + if arguments is None: + raise ValueError("arguments is required for InferenceComponentStep.") + validate_step_arguments( + "InferenceComponentStep", + arguments, + service_name="sagemaker", + operation_name="CreateInferenceComponent", + ) + self._arguments = arguments + self._properties = Properties( + step_name=name, step=self, shape_name="DescribeInferenceComponentOutput" + ) + + @property + def arguments(self) -> RequestType: + """The ``Arguments`` block for the Create/Update InferenceComponent call.""" + validate_step_arguments( + "InferenceComponentStep", + self._arguments, + service_name="sagemaker", + operation_name="CreateInferenceComponent", + ) + return self._arguments + + @property + def properties(self): + """A ``Properties`` object shaped like ``DescribeInferenceComponentOutput``.""" + return self._properties diff --git a/sagemaker-mlops/src/sagemaker/mlops/workflow/lineage_step.py b/sagemaker-mlops/src/sagemaker/mlops/workflow/lineage_step.py new file mode 100644 index 0000000000..20b053d9e1 --- /dev/null +++ b/sagemaker-mlops/src/sagemaker/mlops/workflow/lineage_step.py @@ -0,0 +1,117 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Step definition for SageMaker Lineage tracking in Pipelines. + +Design note: the ``Arguments`` block is a structure of — four optional +lists: + +* ``Actions`` — list of ``CreateActionRequest`` shapes +* ``Artifacts`` — list of ``CreateArtifactRequest`` shapes +* ``Contexts`` — list of ``CreateContextRequest`` shapes +* ``Associations`` — list of ``LineageAssociation`` shapes + (``Source``/``Destination``/``AssociationType``) + +The SDK validates that the ``arguments`` dict contains only these four +top-level keys (at least one required) and forwards it to the service. +""" + +from __future__ import absolute_import + +from typing import Any, Dict, List, Optional, Union + +from sagemaker.core.helper.pipeline_variable import RequestType +from sagemaker.core.workflow.properties import Properties + +from sagemaker.mlops.workflow.step_collections import StepCollection +from sagemaker.mlops.workflow.steps import Step, StepTypeEnum + + +class LineageStep(Step): + """Creates and associates lineage entities in SageMaker's lineage system. + + Wraps SageMaker's ``CreateAction``/``CreateArtifact``/``CreateContext`` + and lineage ``AddAssociation`` APIs. A single step may create + multiple entities of any of the four types (Actions, Artifacts, + Contexts, Associations). Property references use + ``Steps..ActionArns['']``, + ``Steps..ArtifactArns['']``, + ``Steps..ContextArns['']``, and + ``Steps..Associations``. + """ + + def __init__( + self, + name: str, + arguments: Dict[str, Any], + display_name: Optional[str] = None, + description: Optional[str] = None, + depends_on: Optional[List[Union[str, Step, StepCollection]]] = None, + ): + """Construct a ``LineageStep``. + + Args: + name (str): The name of the step. + arguments (Dict[str, Any]): The ``Arguments`` block. Recognized + top-level keys: ``Actions``, ``Artifacts``, ``Contexts``, + ``Associations`` — each is a list of dicts conforming to + the corresponding SageMaker API shape (or the service's + ``LineageAssociation`` for ``Associations``). At least + one of the four keys must be present. + display_name (str): Optional display name. + description (str): Optional description. + depends_on (List[Union[str, Step, StepCollection]]): Optional + explicit step dependencies. + + Raises: + ValueError: If ``arguments`` is None or contains none of the + recognized keys. + """ + super().__init__( + name=name, + display_name=display_name, + description=description, + step_type=StepTypeEnum.LINEAGE, + depends_on=depends_on, + ) + if arguments is None: + raise ValueError("arguments is required for LineageStep.") + if not isinstance(arguments, dict) or not arguments: + raise ValueError("LineageStep: arguments must be a non-empty dict.") + recognized = {"Actions", "Artifacts", "Contexts", "Associations"} + if not recognized & set(arguments.keys()): + raise ValueError( + "LineageStep.arguments must contain at least one of: " + + ", ".join(sorted(recognized)) + ) + unknown = sorted(set(arguments) - recognized) + if unknown: + raise ValueError( + f"LineageStep: unknown argument field(s) {unknown}. " + "Allowed top-level fields: " + ", ".join(sorted(recognized)) + "." + ) + self._arguments = arguments + + root = Properties(step_name=name, step=self) + for field in ("ActionArns", "ArtifactArns", "ContextArns", "Associations"): + root.__dict__[field] = Properties(step_name=name, path=field) + self._properties = root + + @property + def arguments(self) -> RequestType: + """The ``Arguments`` block describing lineage entities and associations.""" + return self._arguments + + @property + def properties(self): + """Exposes ``ActionArns``, ``ArtifactArns``, ``ContextArns``, ``Associations``.""" + return self._properties diff --git a/sagemaker-mlops/src/sagemaker/mlops/workflow/steps.py b/sagemaker-mlops/src/sagemaker/mlops/workflow/steps.py index 76e90a5309..60b7420844 100644 --- a/sagemaker-mlops/src/sagemaker/mlops/workflow/steps.py +++ b/sagemaker-mlops/src/sagemaker/mlops/workflow/steps.py @@ -62,6 +62,10 @@ class StepTypeEnum(Enum): EMR_SERVERLESS = "EMRServerless" FAIL = "Fail" AUTOML = "AutoML" + ENDPOINT_CONFIG = "EndpointConfig" + ENDPOINT = "Endpoint" + INFERENCE_COMPONENT = "InferenceComponent" + LINEAGE = "Lineage" class Step(Entity): diff --git a/sagemaker-mlops/tests/integ/workflow/test_lineage_step.py b/sagemaker-mlops/tests/integ/workflow/test_lineage_step.py new file mode 100644 index 0000000000..3d2d30efe2 --- /dev/null +++ b/sagemaker-mlops/tests/integ/workflow/test_lineage_step.py @@ -0,0 +1,145 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Integration test for the LineageStep. + +Creates a pipeline containing a single ``LineageStep`` that records a +SageMaker lineage Action, executes it end-to-end against the real +service, and asserts the execution reaches ``Succeeded``. Cleans up the +Action, the pipeline, and the S3 pipeline definition artifact. + +Requires the execution role to have ``sagemaker:CreateAction`` (and +related lineage permissions). ``SageMakerRole`` — the standard fixture +role used across the SDK's integ tests — has broad SageMaker access and +satisfies this requirement. + +This test represents the SDK-side end-to-end validation of the +inference and lineage step family. See ``endpoint_step.py`` and +``inference_component_step.py`` for the other step types; those are not +integ-tested here because they provision paid resources +(Endpoint/InferenceComponent). +""" + +from __future__ import absolute_import + +import time +import uuid + +import pytest + +from sagemaker.core.helper.session_helper import Session, get_execution_role +from sagemaker.core.workflow.pipeline_context import PipelineSession +from sagemaker.mlops.workflow.lineage_step import LineageStep +from sagemaker.mlops.workflow.pipeline import Pipeline + + +@pytest.fixture +def sagemaker_session(): + return Session() + + +@pytest.fixture +def pipeline_session(): + return PipelineSession() + + +@pytest.fixture +def role(): + return get_execution_role() + + +def test_lineage_step_execute_end_to_end(sagemaker_session, pipeline_session, role): + """Full end-to-end run of a LineageStep pipeline against the real service. + + Builds a pipeline with a single ``LineageStep`` that creates one + lineage ``Action``. Verifies the pipeline execution succeeds and the + server-reported step metadata contains the created action ARN. + """ + stamp = uuid.uuid4().hex[:8] + action_name = f"lineage-integ-{stamp}" + pipeline_name = f"integ-lineage-{stamp}" + + step = LineageStep( + name="RecordLineage", + arguments={ + "Actions": [ + { + "ActionName": action_name, + "ActionType": "ModelTraining", + "Status": "Completed", + "Source": { + "SourceUri": f"s3://lineage-integ-test/{stamp}/model.tar.gz", + "SourceType": "MODEL", + }, + "Description": "Lineage integ test action", + } + ] + }, + ) + pipeline = Pipeline( + name=pipeline_name, + steps=[step], + sagemaker_session=pipeline_session, + ) + + try: + pipeline.upsert(role_arn=role) + execution = pipeline.start() + + # LineageStep is metadata-only; execution completes quickly. Poll + # up to 5 minutes to give the service plenty of headroom under load. + timeout = 300 + start_time = time.time() + final_status = None + while time.time() - start_time < timeout: + execution_desc = execution.describe() + status = execution_desc["PipelineExecutionStatus"] + if status in ("Succeeded", "Failed", "Stopped"): + final_status = status + break + time.sleep(10) + + if final_status != "Succeeded": + steps = sagemaker_session.sagemaker_client.list_pipeline_execution_steps( + PipelineExecutionArn=execution.arn, + )["PipelineExecutionSteps"] + failure_details = "\n".join( + f"{s['StepName']}: {s.get('FailureReason', 'no reason')}" + for s in steps + if s.get("StepStatus") == "Failed" + ) + pytest.fail(f"Pipeline execution status={final_status}. Details:\n{failure_details}") + + # Verify the step metadata reports the created action ARN. + steps = sagemaker_session.sagemaker_client.list_pipeline_execution_steps( + PipelineExecutionArn=execution.arn, + )["PipelineExecutionSteps"] + lineage_step = next(s for s in steps if s["StepName"] == "RecordLineage") + assert lineage_step["StepStatus"] == "Succeeded" + metadata = lineage_step.get("Metadata", {}) + action_arns = metadata.get("Lineage", {}).get("ActionArns", {}) + assert ( + action_name in action_arns + ), f"expected {action_name} in ActionArns, got: {action_arns}" + assert action_arns[action_name].endswith(f":action/{action_name}") + + finally: + # Delete the lineage Action. + try: + sagemaker_session.sagemaker_client.delete_action(ActionName=action_name) + except Exception: + pass + # Delete the pipeline. + try: + sagemaker_session.sagemaker_client.delete_pipeline(PipelineName=pipeline_name) + except Exception: + pass diff --git a/sagemaker-mlops/tests/unit/workflow/test_inference_lineage_steps.py b/sagemaker-mlops/tests/unit/workflow/test_inference_lineage_steps.py new file mode 100644 index 0000000000..af32493e3c --- /dev/null +++ b/sagemaker-mlops/tests/unit/workflow/test_inference_lineage_steps.py @@ -0,0 +1,332 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Unit tests for the inference and lineage pipeline step types. + +These steps use a passthrough ``arguments: Dict[str, Any]`` API, +mirroring ``LambdaStep``/``CallbackStep``. Top-level argument keys are validated +client-side against the corresponding public AWS API input shape +(botocore service model), and fields known to be rejected by SageMaker +Pipelines fail fast at construction. Values are not validated -- they +may be pipeline variables resolved at compile time. Full schema +validation remains server-side. +""" + +from __future__ import absolute_import + +import pytest + +from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep +from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep +from sagemaker.mlops.workflow.lineage_step import LineageStep +from sagemaker.mlops.workflow.steps import CacheConfig, StepTypeEnum + +# ---------- EndpointConfigStep ---------- + + +def test_endpoint_config_step_basic(): + step = EndpointConfigStep( + name="Cfg", + arguments={ + "EndpointConfigName": "MyCfg", + "ProductionVariants": [ + { + "VariantName": "AllTraffic", + "ModelName": "m", + "InstanceType": "ml.m5.large", + "InitialInstanceCount": 1, + } + ], + }, + ) + assert step.step_type == StepTypeEnum.ENDPOINT_CONFIG + assert step.arguments["EndpointConfigName"] == "MyCfg" + + +def test_endpoint_config_step_to_request_includes_cache_and_retry(): + step = EndpointConfigStep( + name="Cfg", + arguments={"EndpointConfigName": "MyCfg", "ProductionVariants": []}, + display_name="Create Config", + description="desc", + cache_config=CacheConfig(enable_caching=True, expire_after="P30D"), + ) + req = step.to_request() + assert req["Type"] == "EndpointConfig" + assert req["DisplayName"] == "Create Config" + assert req["Description"] == "desc" + assert req["CacheConfig"] == {"Enabled": True, "ExpireAfter": "P30D"} + + +def test_endpoint_config_step_accepts_full_api_surface(): + """User can pass any CreateEndpointConfigInput field (except the ones + the service excludes — that's a server-side rejection, not client-side).""" + step = EndpointConfigStep( + name="Cfg", + arguments={ + "EndpointConfigName": "MyCfg", + "ProductionVariants": [], + "KmsKeyId": "arn:aws:kms:...", + "ExecutionRoleArn": "arn:aws:iam:...", + "AsyncInferenceConfig": {"OutputConfig": {"S3OutputPath": "s3://x/"}}, + "VpcConfig": {"SecurityGroupIds": ["sg-0"], "Subnets": ["subnet-0"]}, + "EnableNetworkIsolation": False, + "ShadowProductionVariants": [], + }, + ) + args = step.arguments + assert args["KmsKeyId"] == "arn:aws:kms:..." + assert args["ExecutionRoleArn"] == "arn:aws:iam:..." + assert "OutputConfig" in args["AsyncInferenceConfig"] + + +def test_endpoint_config_step_requires_arguments(): + with pytest.raises(ValueError): + EndpointConfigStep(name="Cfg", arguments=None) + + +# ---------- EndpointStep ---------- + + +def test_endpoint_step_basic(): + step = EndpointStep( + name="Deploy", + arguments={"EndpointName": "ep", "EndpointConfigName": "cfg"}, + ) + assert step.step_type == StepTypeEnum.ENDPOINT + assert step.arguments == {"EndpointName": "ep", "EndpointConfigName": "cfg"} + + +def test_endpoint_step_cache_config(): + step = EndpointStep( + name="Deploy", + arguments={"EndpointName": "ep", "EndpointConfigName": "cfg"}, + cache_config=CacheConfig(enable_caching=True), + ) + req = step.to_request() + assert req["Type"] == "Endpoint" + assert req["CacheConfig"] == {"Enabled": True} + + +def test_endpoint_step_rejects_retry_policies_kwarg(): + """EndpointStep is not retryable — constructor must not accept retry_policies.""" + with pytest.raises(TypeError): + EndpointStep( + name="Deploy", + arguments={"EndpointName": "ep", "EndpointConfigName": "cfg"}, + retry_policies=[], + ) + + +# ---------- InferenceComponentStep ---------- + + +def test_inference_component_step_basic(): + step = InferenceComponentStep( + name="IC", + arguments={ + "InferenceComponentName": "ic", + "EndpointName": "ep", + "VariantName": "v", + "Specification": { + "ModelName": "m", + "ComputeResourceRequirements": { + "MinMemoryRequiredInMb": 1024, + "NumberOfCpuCoresRequired": 2.0, + }, + }, + "RuntimeConfig": {"CopyCount": 1}, + }, + ) + assert step.step_type == StepTypeEnum.INFERENCE_COMPONENT + assert step.arguments["Specification"]["ModelName"] == "m" + + +def test_inference_component_step_rejects_retry_policies_kwarg(): + with pytest.raises(TypeError): + InferenceComponentStep( + name="IC", + arguments={}, + retry_policies=[], + ) + + +# ---------- LineageStep ---------- + + +def test_lineage_step_basic(): + step = LineageStep( + name="Rec", + arguments={ + "Actions": [ + { + "ActionName": "a1", + "ActionType": "ModelTraining", + "Status": "Completed", + } + ], + "Artifacts": [ + { + "ArtifactName": "art1", + "ArtifactType": "Model", + "Source": {"SourceUri": "s3://x/y"}, + } + ], + "Associations": [ + { + "Source": {"Name": "a1", "Type": "Action"}, + "Destination": {"Name": "art1", "Type": "Artifact"}, + "AssociationType": "Produced", + } + ], + }, + ) + assert step.step_type == StepTypeEnum.LINEAGE + assert len(step.arguments["Actions"]) == 1 + assert len(step.arguments["Associations"]) == 1 + + +def test_lineage_step_partial_arguments(): + step = LineageStep( + name="Rec", + arguments={"Actions": [{"ActionName": "a", "ActionType": "T", "Status": "Completed"}]}, + ) + assert "Actions" in step.arguments + assert "Artifacts" not in step.arguments + + +def test_lineage_step_requires_at_least_one_recognized_key(): + with pytest.raises(ValueError): + LineageStep(name="Rec", arguments={}) + with pytest.raises(ValueError): + LineageStep(name="Rec", arguments={"Bogus": []}) + + +def test_lineage_step_properties(): + step = LineageStep(name="Rec", arguments={"Actions": []}) + for field in ("ActionArns", "ArtifactArns", "ContextArns", "Associations"): + assert hasattr(step.properties, field) + + +# ---------- Cross-cutting ---------- + + +def test_all_steps_importable_from_init(): + from sagemaker.mlops.workflow import ( # noqa: F401 + EndpointConfigStep, + EndpointStep, + InferenceComponentStep, + LineageStep, + ) + + +def test_step_type_enum_values(): + assert StepTypeEnum.ENDPOINT_CONFIG.value == "EndpointConfig" + assert StepTypeEnum.ENDPOINT.value == "Endpoint" + assert StepTypeEnum.INFERENCE_COMPONENT.value == "InferenceComponent" + assert StepTypeEnum.LINEAGE.value == "Lineage" + + +def test_depends_on_accepts_string_list(): + step = EndpointStep( + name="Deploy", + arguments={"EndpointName": "ep", "EndpointConfigName": "cfg"}, + depends_on=["Prev"], + ) + req = step.to_request() + assert req["DependsOn"] == ["Prev"] + + +# ---------- Client-side argument validation ---------- + + +def test_endpoint_config_step_rejects_unsupported_fields(): + """DataCaptureConfig and ExplainerConfig exist in the public API but + are rejected by SageMaker Pipelines -- fail fast with a clear error.""" + for field in ("DataCaptureConfig", "ExplainerConfig"): + with pytest.raises(ValueError, match=field): + EndpointConfigStep( + name="Cfg", + arguments={ + "EndpointConfigName": "cfg", + "ProductionVariants": [], + field: {}, + }, + ) + + +def test_endpoint_step_rejects_unsupported_deployment_config(): + with pytest.raises(ValueError, match="DeploymentConfig"): + EndpointStep( + name="Deploy", + arguments={ + "EndpointName": "ep", + "EndpointConfigName": "cfg", + "DeploymentConfig": {}, + }, + ) + + +def test_unknown_argument_key_rejected(): + """Keys outside the operation's input shape fail fast at construction.""" + with pytest.raises(ValueError, match="Bogus"): + EndpointConfigStep( + name="Cfg", + arguments={"EndpointConfigName": "cfg", "Bogus": 1}, + ) + with pytest.raises(ValueError, match="Bogus"): + InferenceComponentStep( + name="IC", + arguments={"InferenceComponentName": "ic", "Bogus": 1}, + ) + + +def test_empty_arguments_rejected(): + for cls, valid_key in ( + (EndpointConfigStep, "EndpointConfigName"), + (EndpointStep, "EndpointName"), + (InferenceComponentStep, "InferenceComponentName"), + ): + with pytest.raises(ValueError): + cls(name="x", arguments={}) + # sanity: a single valid key constructs fine + assert cls(name="x", arguments={valid_key: "v"}).arguments == {valid_key: "v"} + + +def test_pipeline_variable_values_pass_validation(): + """Only top-level keys are validated -- values may be pipeline + variables (Get expressions) at any position.""" + step = EndpointStep( + name="Deploy", + arguments={ + "EndpointName": {"Get": "Parameters.EndpointName"}, + "EndpointConfigName": {"Get": "Steps.Cfg.EndpointConfigName"}, + }, + ) + assert step.arguments["EndpointName"] == {"Get": "Parameters.EndpointName"} + + +def test_post_construction_mutation_caught_at_serialization(): + """Injecting an unsupported field after construction is caught when + the arguments property is read (i.e., at pipeline serialization).""" + step = EndpointConfigStep( + name="Cfg", + arguments={"EndpointConfigName": "cfg", "ProductionVariants": []}, + ) + step._arguments["DataCaptureConfig"] = {} + with pytest.raises(ValueError, match="DataCaptureConfig"): + _ = step.arguments + + +def test_lineage_step_rejects_unknown_keys_alongside_recognized(): + with pytest.raises(ValueError, match="Bogus"): + LineageStep(name="Rec", arguments={"Actions": [], "Bogus": []})