import json import time from typing import Any, TYPE_CHECKING import boto3 from aws_testing_utils import config if TYPE_CHECKING: from mypy_boto3_stepfunctions import SFNClient class StepFunctionHandler: """Handles interactions with AWS Step Functions.""" client: 'SFNClient' function_name: str arn: str last_execution: str def __init__(self, function_name: str) -> None: session = boto3.session.Session() self.client = session.client( service_name='stepfunctions', region_name=config.AWS_REGION ) self.function_name = function_name self.arn = self.set_machine_arn(function_name) self.last_execution = '' def set_last_execution(self) -> None: """Snapshots the current last execution ARN for later comparison.""" self.last_execution = self.get_last_execution() def get_last_execution(self) -> str: """Returns the ARN of the most recent execution.""" response = self.client.list_executions(stateMachineArn=self.arn, maxResults=1) execution_arn = response['executions'][0]['executionArn'] return execution_arn def set_machine_arn(self, function_name: str) -> str: """Resolves and returns the state machine ARN for the given function name.""" paginator = self.client.get_paginator('list_state_machines') for page in paginator.paginate(): for state_machine in page.get('stateMachines', []): if state_machine['name'] == function_name: return state_machine['stateMachineArn'] raise ValueError(f'Step function {function_name!r} not found') def assert_machine_ran(self) -> bool: """Asserts a new execution has started since set_last_execution() was called.""" tries = 0 max_tries = 10 while tries < max_tries: latest_arn = self.get_last_execution() if latest_arn != self.last_execution: return True tries += 1 time.sleep(5) raise AssertionError(f'Step function {self.function_name} was never triggered') def execute( self, input_data: dict[str, Any], assert_success: bool = True, timeout: int = 120, execution_name: str | None = None, **kwargs: Any, ) -> None: """Starts a step function execution. Args: input_data: Input payload for the execution. assert_success: If True, polls until SUCCEEDED or raises on failure/timeout. timeout: Maximum seconds to wait for completion. execution_name: Optional execution name. kwargs: Extra keyword arguments forwarded straight to boto3's ``start_execution`` (e.g. ``traceHeader``). """ # Reserved keys are applied after kwargs so they always win over any # collision — callers can't override the resolved ARN or input, and an # explicit execution_name takes precedence over a name in kwargs. start_args: dict[str, Any] = { **kwargs, 'stateMachineArn': self.arn, 'input': json.dumps(input_data), } if execution_name: start_args['name'] = execution_name response = self.client.start_execution(**start_args) if not assert_success: return execution_arn = response['executionArn'] self.assert_execution_success(execution_arn, timeout) def assert_execution_success(self, execution_arn: str, max_wait: int) -> None: """Polls until the execution succeeds, or raises if it fails or times out. Args: execution_arn: ARN of the execution to poll. max_wait: Maximum seconds to wait. """ tries = 0 max_tries = max_wait while tries < max_tries: status = self.client.describe_execution(executionArn=execution_arn)[ 'status' ] if status == 'RUNNING': tries += 1 time.sleep(1) continue assert status == 'SUCCEEDED', f'Execution status is {status}' return raise TimeoutError( f'State machine execution {execution_arn} still running after {max_wait} seconds' )