diff --git a/.github/workflows/pythonbuild.yml b/.github/workflows/pythonbuild.yml index 261a09f1c3..5a8c2eacb7 100644 --- a/.github/workflows/pythonbuild.yml +++ b/.github/workflows/pythonbuild.yml @@ -294,6 +294,7 @@ jobs: - flytekit-async-fsspec - flytekit-aws-athena - flytekit-aws-batch + - flytekit-aws-emr-serverless - flytekit-aws-sagemaker - flytekit-bigquery - flytekit-comet-ml diff --git a/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/boto_handler.py b/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/boto_handler.py index 0cfb82e2e0..18bfa9f59a 100644 --- a/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/boto_handler.py +++ b/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/boto_handler.py @@ -287,6 +287,32 @@ async def ensure_application_started( ) await asyncio.sleep(poll_interval_seconds) + async def start_application_if_needed(self, application_id: str) -> bool: + """Request application startup without waiting for the transition. + + Returns ``True`` when the application is already ``STARTED``. For + ``CREATED`` or ``STOPPED`` applications, this sends ``StartApplication`` + and returns ``False``. Other non-terminal transitional states also + return ``False`` so the connector can continue the startup from its + polling path instead of blocking the CreateTask RPC. + """ + app = await self.get_application(application_id) + state = app.get("state", "") + + if state == _APP_STARTED: + return True + + if state in _APP_TERMINAL: + raise RuntimeError(f"Application {application_id} is in terminal state '{state}' and cannot be started") + + if state in _APP_NEEDS_START: + logger.info("Application %s is in state '%s', sending StartApplication request", application_id, state) + await self._call("start_application", applicationId=application_id) + else: + logger.info("Application %s is transitioning in state '%s'", application_id, state) + + return False + # ------------------------------------------------------------------ # Job management # ------------------------------------------------------------------ @@ -301,6 +327,7 @@ async def start_job_run( execution_timeout_minutes: int = 60, name: Optional[str] = None, retry_policy: Optional[Dict[str, Any]] = None, + client_token: Optional[str] = None, ) -> str: logger.info( "StartJobRun: applicationId=%s, name=%s, timeout=%dm", @@ -324,6 +351,8 @@ async def start_job_run( if retry_policy: params["retryPolicy"] = retry_policy logger.debug("StartJobRun: retryPolicy=%s", retry_policy) + if client_token: + params["clientToken"] = client_token logger.debug("StartJobRun: jobDriver type=%s", list(job_driver.keys())) resp = await self._call("start_job_run", **params) diff --git a/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/connector.py b/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/connector.py index 52c41a9f8b..c5ea841afa 100644 --- a/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/connector.py +++ b/plugins/flytekit-aws-emr-serverless/flytekitplugins/awsemrserverless/connector.py @@ -9,6 +9,7 @@ import logging import os import re +import uuid from dataclasses import dataclass from pathlib import Path from typing import Any, Dict, Optional @@ -75,6 +76,8 @@ class EMRServerlessJobMetadata(ResourceMeta): job_run_id: str region: str created_application: bool = False + is_script_mode: bool = False + pending_job_request: Optional[Dict[str, Any]] = None class EMRServerlessConnector(AsyncConnectorBase): @@ -655,9 +658,6 @@ async def create( elif not created_application: logger.debug("sync_image is disabled, skipping image sync for %s", application_id) - logger.info("Ensuring application %s is in STARTED state", application_id) - await handler.ensure_application_started(application_id) - # --- Build job driver --- if config.is_script_mode: logger.info("Building job driver in script mode") @@ -704,6 +704,35 @@ async def create( list(effective_config_overrides.keys()), ) + client_token = uuid.uuid4().hex + job_request = { + "execution_role_arn": config.execution_role_arn, + "job_driver": job_driver, + "configuration_overrides": effective_config_overrides, + "tags": self._merge_tags(config.tags), + "execution_timeout_minutes": config.execution_timeout_minutes, + "name": job_name, + "retry_policy": config.retry_policy, + "client_token": client_token, + } + region = config.region or handler.client.meta.region_name + + logger.info("Ensuring application %s startup has been requested", application_id) + application_started = await handler.start_application_if_needed(application_id) + if not application_started: + logger.info( + "Application %s is still starting; deferring job submission to get()", + application_id, + ) + return EMRServerlessJobMetadata( + application_id=application_id, + job_run_id="", + region=region, + created_application=created_application, + is_script_mode=config.is_script_mode, + pending_job_request=job_request, + ) + logger.info( "Submitting job run: application=%s, job_name=%s, execution_role=%s, timeout=%dm", application_id, @@ -713,16 +742,9 @@ async def create( ) job_run_id = await handler.start_job_run( application_id=application_id, - execution_role_arn=config.execution_role_arn, - job_driver=job_driver, - configuration_overrides=effective_config_overrides, - tags=self._merge_tags(config.tags), - execution_timeout_minutes=config.execution_timeout_minutes, - name=job_name, - retry_policy=config.retry_policy, + **job_request, ) - region = config.region or handler.client.meta.region_name logger.info( "Job submitted successfully: application=%s, job_run_id=%s, region=%s, created_application=%s", application_id, @@ -736,6 +758,7 @@ async def create( job_run_id=job_run_id, region=region, created_application=created_application, + is_script_mode=config.is_script_mode, ) async def get( @@ -750,22 +773,50 @@ async def get( resource_meta.region, ) handler = self._get_handler(resource_meta.region) + job_run_id = resource_meta.job_run_id + + if not job_run_id: + if not resource_meta.pending_job_request: + return Resource( + phase=TaskExecution.FAILED, + message="Job submission metadata is missing", + ) + + try: + application_started = await handler.start_application_if_needed(resource_meta.application_id) + except RuntimeError as e: + return Resource(phase=TaskExecution.FAILED, message=str(e)) + + if not application_started: + return Resource( + phase=TaskExecution.RUNNING, + message=f"EMR Serverless application {resource_meta.application_id} is starting", + ) + + logger.info( + "Application %s is STARTED; submitting deferred job", + resource_meta.application_id, + ) + job_run_id = await handler.start_job_run( + application_id=resource_meta.application_id, + **resource_meta.pending_job_request, + ) try: job = await handler.get_job_run( application_id=resource_meta.application_id, - job_run_id=resource_meta.job_run_id, + job_run_id=job_run_id, ) except Exception as e: logger.warning( "Failed to retrieve job %s on application %s: %s", - resource_meta.job_run_id, + job_run_id, resource_meta.application_id, e, ) return Resource( phase=TaskExecution.FAILED, - message=f"Job not found: {resource_meta.job_run_id}", + message=f"Job not found: {job_run_id}", ) state = job.get("state", "UNKNOWN") @@ -778,20 +829,22 @@ async def get( logger.info( "Job %s status: state=%s, phase=%s", - resource_meta.job_run_id, + job_run_id, state, phase, ) - log_links = self._get_log_links(resource_meta) + log_links = self._get_log_links(resource_meta, job_run_id) + outputs = LiteralMap(literals={}) if phase == TaskExecution.SUCCEEDED and resource_meta.is_script_mode else None - return Resource(phase=phase, message=message, log_links=log_links) + return Resource(phase=phase, message=message, log_links=log_links, outputs=outputs) - def _get_log_links(self, resource_meta: EMRServerlessJobMetadata) -> list: + def _get_log_links(self, resource_meta: EMRServerlessJobMetadata, job_run_id: Optional[str] = None) -> list: region = resource_meta.region or "us-east-1" + resolved_job_run_id = job_run_id or resource_meta.job_run_id console_url = ( f"https://{region}.console.aws.amazon.com/emr/home?region={region}" - f"#/serverless/{resource_meta.application_id}/jobs/{resource_meta.job_run_id}" + f"#/serverless/{resource_meta.application_id}/jobs/{resolved_job_run_id}" ) return [TaskLog(uri=console_url, name="EMR Serverless Console").to_flyte_idl()] @@ -807,16 +860,34 @@ async def delete( resource_meta.region, ) handler = self._get_handler(resource_meta.region) + job_run_id = resource_meta.job_run_id try: + if not job_run_id and resource_meta.pending_job_request: + application_started = await handler.start_application_if_needed(resource_meta.application_id) + if not application_started: + logger.info( + "Application %s is still starting; no submitted job to cancel", + resource_meta.application_id, + ) + return + job_run_id = await handler.start_job_run( + application_id=resource_meta.application_id, + **resource_meta.pending_job_request, + ) + + if not job_run_id: + logger.info("No submitted job to cancel for application %s", resource_meta.application_id) + return + await handler.cancel_job_run( application_id=resource_meta.application_id, - job_run_id=resource_meta.job_run_id, + job_run_id=job_run_id, ) - logger.info("Delete completed for job %s", resource_meta.job_run_id) + logger.info("Delete completed for job %s", job_run_id) except Exception as e: logger.warning( "Failed to cancel job %s on application %s: %s", - resource_meta.job_run_id, + job_run_id, resource_meta.application_id, e, ) diff --git a/plugins/flytekit-aws-emr-serverless/tests/test_boto_handler.py b/plugins/flytekit-aws-emr-serverless/tests/test_boto_handler.py index cf2c75636e..7a702f5474 100644 --- a/plugins/flytekit-aws-emr-serverless/tests/test_boto_handler.py +++ b/plugins/flytekit-aws-emr-serverless/tests/test_boto_handler.py @@ -91,6 +91,7 @@ async def test_start_job_run_with_all_options(self, mock_call): tags=tags, execution_timeout_minutes=120, name="test-job", + client_token="stable-token", ) call_kwargs = mock_call.call_args.kwargs @@ -98,6 +99,7 @@ async def test_start_job_run_with_all_options(self, mock_call): assert call_kwargs["tags"] == tags assert call_kwargs["executionTimeoutMinutes"] == 120 assert call_kwargs["name"] == "test-job" + assert call_kwargs["clientToken"] == "stable-token" @pytest.mark.asyncio async def test_start_job_run_with_retry_policy(self, mock_call): @@ -334,3 +336,49 @@ async def test_raises_on_terminal_state(self, mock_call): handler = EMRServerlessHandler() with pytest.raises(RuntimeError, match="terminal state"): await handler.ensure_application_started("app-1") + + +class TestEMRServerlessHandlerStartApplicationIfNeeded: + @pytest.mark.asyncio + async def test_returns_true_when_started(self, mock_call): + mock_call.return_value = {"application": {"applicationId": "app-1", "state": "STARTED"}} + + handler = EMRServerlessHandler() + result = await handler.start_application_if_needed("app-1") + + assert result is True + mock_call.assert_awaited_once_with("get_application", applicationId="app-1") + + @pytest.mark.asyncio + async def test_requests_start_without_waiting(self, mock_call): + mock_call.side_effect = [ + {"application": {"applicationId": "app-1", "state": "STOPPED"}}, + {}, + ] + + handler = EMRServerlessHandler() + result = await handler.start_application_if_needed("app-1") + + assert result is False + assert [call.args[0] for call in mock_call.await_args_list] == [ + "get_application", + "start_application", + ] + + @pytest.mark.asyncio + async def test_returns_false_for_transitioning_application(self, mock_call): + mock_call.return_value = {"application": {"applicationId": "app-1", "state": "CREATING"}} + + handler = EMRServerlessHandler() + result = await handler.start_application_if_needed("app-1") + + assert result is False + mock_call.assert_awaited_once_with("get_application", applicationId="app-1") + + @pytest.mark.asyncio + async def test_raises_on_terminal_state(self, mock_call): + mock_call.return_value = {"application": {"applicationId": "app-1", "state": "TERMINATED"}} + + handler = EMRServerlessHandler() + with pytest.raises(RuntimeError, match="terminal state"): + await handler.start_application_if_needed("app-1") diff --git a/plugins/flytekit-aws-emr-serverless/tests/test_connector.py b/plugins/flytekit-aws-emr-serverless/tests/test_connector.py index fb0ba4d13e..5d29e74f94 100644 --- a/plugins/flytekit-aws-emr-serverless/tests/test_connector.py +++ b/plugins/flytekit-aws-emr-serverless/tests/test_connector.py @@ -25,6 +25,7 @@ def _make_handler(**overrides) -> AsyncMock: """Build a mock EMRServerlessHandler with sensible defaults.""" handler = AsyncMock() handler.ensure_application_started = AsyncMock() + handler.start_application_if_needed = AsyncMock(return_value=True) handler.start_job_run = AsyncMock(return_value="job-123") handler.create_application = AsyncMock(return_value="new-app-123") handler.get_application = AsyncMock(return_value={ @@ -51,6 +52,8 @@ def test_metadata_creation(self): assert metadata.job_run_id == "job-456" assert metadata.region == "us-east-1" assert metadata.created_application is False + assert metadata.is_script_mode is False + assert metadata.pending_job_request is None def test_metadata_with_created_app(self): metadata = EMRServerlessJobMetadata( @@ -61,6 +64,23 @@ def test_metadata_with_created_app(self): ) assert metadata.created_application is True + def test_pending_metadata_round_trip(self): + metadata = EMRServerlessJobMetadata( + application_id="app-123", + job_run_id="", + region="us-east-1", + is_script_mode=True, + pending_job_request={ + "execution_role_arn": "arn:aws:iam::123456789012:role/Role", + "job_driver": {"sparkSubmit": {"entryPoint": "s3://bucket/main.py"}}, + "client_token": "stable-token", + }, + ) + + decoded = EMRServerlessJobMetadata.decode(metadata.encode()) + + assert decoded == metadata + class TestEMRServerlessConnector: def test_connector_initialization(self): @@ -84,9 +104,32 @@ async def test_create_with_existing_application(self, sample_config): assert isinstance(metadata, EMRServerlessJobMetadata) assert metadata.application_id == sample_config.application_id assert metadata.job_run_id == "job-123" - mock_handler.ensure_application_started.assert_called_once() + assert metadata.is_script_mode is True + mock_handler.start_application_if_needed.assert_called_once() mock_handler.start_job_run.assert_called_once() + @pytest.mark.asyncio + async def test_create_defers_job_submission_while_application_starts(self, sample_config): + connector = EMRServerlessConnector() + mock_handler = _make_handler() + mock_handler.start_application_if_needed = AsyncMock(return_value=False) + + mock_template = MagicMock() + mock_template.custom = sample_config.to_dict() + mock_template.id = MagicMock() + mock_template.id.name = "test-task" + + with patch.object(connector, "_get_handler", return_value=mock_handler): + with patch.object(connector, "_extract_config", return_value=sample_config): + metadata = await connector.create(mock_template) + + assert metadata.application_id == sample_config.application_id + assert metadata.job_run_id == "" + assert metadata.is_script_mode is True + assert metadata.pending_job_request is not None + assert metadata.pending_job_request["client_token"] + mock_handler.start_job_run.assert_not_called() + @pytest.mark.asyncio async def test_create_with_new_application(self, sample_spark_job_driver): """application_name + env var true => create succeeds.""" @@ -760,6 +803,74 @@ async def test_get_successful_job(self, sample_job_metadata): resource = await connector.get(sample_job_metadata) assert resource.phase == TaskExecution.SUCCEEDED + assert resource.outputs is None + + @pytest.mark.asyncio + async def test_get_waits_for_application_before_deferred_submission(self): + connector = EMRServerlessConnector() + metadata = EMRServerlessJobMetadata( + application_id="00f5abc123def456", + job_run_id="", + region="us-east-1", + is_script_mode=True, + pending_job_request={ + "execution_role_arn": "arn:aws:iam::123456789012:role/Role", + "job_driver": {"sparkSubmit": {"entryPoint": "s3://bucket/main.py"}}, + "client_token": "stable-token", + }, + ) + mock_handler = _make_handler() + mock_handler.start_application_if_needed = AsyncMock(return_value=False) + + with patch.object(connector, "_get_handler", return_value=mock_handler): + resource = await connector.get(metadata) + + assert resource.phase == TaskExecution.RUNNING + assert "is starting" in resource.message + mock_handler.start_job_run.assert_not_called() + mock_handler.get_job_run.assert_not_called() + + @pytest.mark.asyncio + async def test_get_submits_deferred_job_idempotently_and_returns_script_outputs(self): + connector = EMRServerlessConnector() + pending_request = { + "execution_role_arn": "arn:aws:iam::123456789012:role/Role", + "job_driver": {"sparkSubmit": {"entryPoint": "s3://bucket/main.py"}}, + "client_token": "stable-token", + } + metadata = EMRServerlessJobMetadata( + application_id="00f5abc123def456", + job_run_id="", + region="us-east-1", + is_script_mode=True, + pending_job_request=pending_request, + ) + mock_handler = _make_handler() + mock_handler.start_job_run = AsyncMock(return_value="deferred-job-123") + mock_handler.get_job_run = AsyncMock( + return_value={ + "state": "SUCCESS", + "stateDetails": "Job completed successfully", + } + ) + + with patch.object(connector, "_get_handler", return_value=mock_handler): + resource = await connector.get(metadata) + + mock_handler.start_job_run.assert_awaited_once_with( + application_id=metadata.application_id, + **pending_request, + ) + mock_handler.get_job_run.assert_awaited_once_with( + application_id=metadata.application_id, + job_run_id="deferred-job-123", + ) + assert resource.phase == TaskExecution.SUCCEEDED + assert resource.outputs is not None + assert resource.outputs.literals == {} + assert "deferred-job-123" in resource.log_links[0].uri + resource_idl = await resource.to_flyte_idl() + assert resource_idl.HasField("outputs") @pytest.mark.asyncio async def test_get_failed_job(self, sample_job_metadata):