diff --git a/kubeflow/spark/backends/kubernetes/backend.py b/kubeflow/spark/backends/kubernetes/backend.py index bdfc2adf5..fde6b6a1b 100644 --- a/kubeflow/spark/backends/kubernetes/backend.py +++ b/kubeflow/spark/backends/kubernetes/backend.py @@ -367,7 +367,7 @@ def _wait_for_session_ready( while True: info = self.get_session(name) - if info.state in (SparkConnectState.READY, SparkConnectState.RUNNING): + if info.state == SparkConnectState.READY: logger.info( "Session ready: %s/%s state=%s serviceName=%s (%.0fs)", self.namespace, @@ -668,6 +668,10 @@ def create_and_connect( 2. Waits for session to become ready 3. Connects to the session and returns SparkSession + If waiting for readiness or connecting fails after the SparkConnect CR has been + created, the session is deleted in a best-effort manner so failed setup does not + leave orphaned cluster resources. + Args: num_executors: Number of executor instances. resources_per_executor: Resource requirements per executor. @@ -703,10 +707,29 @@ def create_and_connect( timeout, ) - info = self._wait_for_session_ready(info.name, timeout=timeout) - logger.info("Session ready, connecting (service_name=%s)", info.service_name) + try: + info = self._wait_for_session_ready(info.name, timeout=timeout) + logger.info("Session ready, connecting (service_name=%s)", info.service_name) + return self.connect(info, connect_timeout=connect_timeout) + except Exception: + self._cleanup_session_on_failure(info.name) + raise - return self.connect(info, connect_timeout=connect_timeout) + def _cleanup_session_on_failure(self, name: str) -> None: + """Best-effort delete of a SparkConnect session after setup/connect failure. + + Args: + name: Name of the SparkConnect session to delete. + """ + try: + self.delete_session(name) + except Exception as cleanup_error: + logger.warning( + "Failed to clean up SparkConnect session %s/%s after setup failure: %s", + self.namespace, + name, + cleanup_error, + ) def get_session_logs( self, diff --git a/kubeflow/spark/backends/kubernetes/backend_test.py b/kubeflow/spark/backends/kubernetes/backend_test.py index f954dc8c8..256c3817f 100644 --- a/kubeflow/spark/backends/kubernetes/backend_test.py +++ b/kubeflow/spark/backends/kubernetes/backend_test.py @@ -853,11 +853,13 @@ def test_create_and_connect(kubernetes_backend, test_case): patch.object( kubernetes_backend, "get_connect_url", return_value=("sc://localhost:15002", None) ), + patch.object(kubernetes_backend, "delete_session") as mock_delete, patch("kubeflow.spark.backends.kubernetes.backend.SparkSession"), ): kubernetes_backend.create_and_connect(options=options) mock_create.assert_called_once() assert mock_create.call_args.kwargs.get("options") == options + mock_delete.assert_not_called() assert test_case.expected_status == SUCCESS @@ -866,6 +868,99 @@ def test_create_and_connect(kubernetes_backend, test_case): print("test execution complete") +@pytest.mark.parametrize( + "test_case", + [ + TestCase( + name="cleanup when wait for ready fails", + expected_status=FAILED, + expected_error=RuntimeError, + config={ + "session_name": "spark-connect-wait-fail", + "fail_at": "wait", + "wait_error": RuntimeError("SparkConnect failed"), + }, + ), + TestCase( + name="cleanup when wait for ready times out", + expected_status=TIMEOUT, + expected_error=TimeoutError, + config={ + "session_name": "spark-connect-wait-timeout", + "fail_at": "wait", + "wait_error": TimeoutError("Timeout waiting for SparkConnect"), + }, + ), + TestCase( + name="cleanup when connect fails", + expected_status=FAILED, + expected_error=RuntimeError, + config={ + "session_name": "spark-connect-connect-fail", + "fail_at": "connect", + "connect_error": RuntimeError("Port-forward failed"), + }, + ), + TestCase( + name="original error preserved when cleanup delete fails", + expected_status=FAILED, + expected_error=RuntimeError, + config={ + "session_name": "spark-connect-cleanup-fail", + "fail_at": "wait", + "wait_error": RuntimeError("SparkConnect failed"), + "delete_error": RuntimeError("delete failed"), + }, + ), + ], +) +def test_create_and_connect_cleans_up_on_failure(kubernetes_backend, test_case): + """Test create_and_connect deletes the session when setup/connect fails.""" + print("Executing test:", test_case.name) + session_name = test_case.config["session_name"] + created_info = SparkConnectInfo( + name=session_name, + namespace=DEFAULT_NAMESPACE, + state=SparkConnectState.PROVISIONING, + service_name="svc", + ) + ready_info = SparkConnectInfo( + name=session_name, + namespace=DEFAULT_NAMESPACE, + state=SparkConnectState.READY, + service_name="svc", + ) + + wait_patch_kwargs: dict = {} + connect_patch_kwargs: dict = {"return_value": Mock()} + if test_case.config["fail_at"] == "wait": + wait_patch_kwargs["side_effect"] = test_case.config["wait_error"] + else: + wait_patch_kwargs["return_value"] = ready_info + connect_patch_kwargs["side_effect"] = test_case.config["connect_error"] + + delete_patch_kwargs: dict = {} + if test_case.config.get("delete_error") is not None: + delete_patch_kwargs["side_effect"] = test_case.config["delete_error"] + + with ( + patch.object(kubernetes_backend, "_create_session", return_value=created_info), + patch.object(kubernetes_backend, "_wait_for_session_ready", **wait_patch_kwargs), + patch.object(kubernetes_backend, "connect", **connect_patch_kwargs), + patch.object(kubernetes_backend, "delete_session", **delete_patch_kwargs) as mock_delete, + ): + with pytest.raises(test_case.expected_error) as exc_info: + kubernetes_backend.create_and_connect() + + if test_case.config["fail_at"] == "wait": + assert exc_info.value is test_case.config["wait_error"] + else: + assert exc_info.value is test_case.config["connect_error"] + mock_delete.assert_called_once_with(session_name) + + print("test execution complete") + + @pytest.mark.parametrize( "test_case", [ diff --git a/kubeflow/spark/backends/kubernetes/utils.py b/kubeflow/spark/backends/kubernetes/utils.py index ca7fdc4b6..57870540e 100644 --- a/kubeflow/spark/backends/kubernetes/utils.py +++ b/kubeflow/spark/backends/kubernetes/utils.py @@ -527,13 +527,13 @@ def get_spark_connect_info_from_cr( if not (spark_connect_cr.metadata and spark_connect_cr.metadata.name): raise ValueError(f"SparkConnect CR is invalid: {spark_connect_cr}") - # Parse state - state = SparkConnectState.PROVISIONING - if spark_connect_cr.status and spark_connect_cr.status.state: + # Parse state. Operator New is ""; use `is not None` so empty string is preserved. + state = SparkConnectState.NEW + if spark_connect_cr.status and spark_connect_cr.status.state is not None: try: state = SparkConnectState(spark_connect_cr.status.state) except ValueError: - state = SparkConnectState.PROVISIONING + state = SparkConnectState.NEW # Extract server status server_status = None diff --git a/kubeflow/spark/backends/kubernetes/utils_test.py b/kubeflow/spark/backends/kubernetes/utils_test.py index 344db527c..b109818c8 100644 --- a/kubeflow/spark/backends/kubernetes/utils_test.py +++ b/kubeflow/spark/backends/kubernetes/utils_test.py @@ -652,19 +652,57 @@ def test_get_spark_connect_info_from_cr( elif test_case.name == "failed status": assert info.state == SparkConnectState.FAILED - elif test_case.name == "running status": - assert info.state == SparkConnectState.RUNNING - assert info.service_name == "run-session-svc" + def test_parse_new_status(self, minimal_spec): + """Parse CR with operator New state (empty string).""" + spark_connect_cr = models.SparkV1alpha1SparkConnect( + metadata=models.IoK8sApimachineryPkgApisMetaV1ObjectMeta( + name="new-session", + namespace="default", + ), + spec=minimal_spec, + status=models.SparkV1alpha1SparkConnectStatus( + state="", + server=models.SparkV1alpha1SparkConnectServerStatus( + podName="new-session-server", + serviceName="new-session-svc", + ), + ), + ) + info = get_spark_connect_info_from_cr(spark_connect_cr) + assert info.state == SparkConnectState.NEW + assert info.service_name == "new-session-svc" + + def test_parse_unknown_status_defaults_to_new(self, minimal_spec): + """Unknown CR states fall back to New.""" + spark_connect_cr = models.SparkV1alpha1SparkConnect( + metadata=models.IoK8sApimachineryPkgApisMetaV1ObjectMeta( + name="unknown-session", + namespace="default", + ), + spec=minimal_spec, + status=models.SparkV1alpha1SparkConnectStatus( + state="Running", + ), + ) + info = get_spark_connect_info_from_cr(spark_connect_cr) + assert info.state == SparkConnectState.NEW elif test_case.name == "empty status": assert info.state == SparkConnectState.PROVISIONING assert info.driver_pod_name is None - else: - with pytest.raises( - test_case.expected_error, - match=test_case.expected_output, - ): + assert info.state == SparkConnectState.NEW + assert info.driver_pod_name is None + + def test_invalid_cr_missing_name_raises_error(self, minimal_spec): + """Test that CR without name in metadata raises ValueError.""" + spark_connect_cr = models.SparkV1alpha1SparkConnect( + metadata=models.IoK8sApimachineryPkgApisMetaV1ObjectMeta( + namespace="default", + ), + spec=minimal_spec, + ) + with pytest.raises(ValueError, match="SparkConnect CR is invalid"): get_spark_connect_info_from_cr(spark_connect_cr) print("test execution complete") diff --git a/kubeflow/spark/types/types.py b/kubeflow/spark/types/types.py index 60c2ffb9a..a66704e7f 100644 --- a/kubeflow/spark/types/types.py +++ b/kubeflow/spark/types/types.py @@ -25,11 +25,15 @@ class SparkConnectState(str, Enum): - """State of a SparkConnect session.""" + """State of a SparkConnect session. + Values match the Spark Operator SparkConnectState model. The operator uses an + empty string for the New state, so avoid truthiness checks on ``status.state``. + """ + + NEW = "" PROVISIONING = "Provisioning" READY = "Ready" - RUNNING = "Running" # Operator may set this when server is up; treated as ready NOT_READY = "NotReady" FAILED = "Failed" diff --git a/kubeflow/spark/types/types_test.py b/kubeflow/spark/types/types_test.py index f32929943..839d783e0 100644 --- a/kubeflow/spark/types/types_test.py +++ b/kubeflow/spark/types/types_test.py @@ -35,9 +35,9 @@ @pytest.mark.parametrize( "state, expected", [ + (SparkConnectState.NEW, ""), (SparkConnectState.PROVISIONING, "Provisioning"), (SparkConnectState.READY, "Ready"), - (SparkConnectState.RUNNING, "Running"), (SparkConnectState.NOT_READY, "NotReady"), (SparkConnectState.FAILED, "Failed"), ], @@ -50,9 +50,9 @@ def test_spark_connect_state_values(state, expected): @pytest.mark.parametrize( "state", [ + SparkConnectState.NEW, SparkConnectState.PROVISIONING, SparkConnectState.READY, - SparkConnectState.RUNNING, SparkConnectState.NOT_READY, SparkConnectState.FAILED, ], diff --git a/test/e2e/spark/test_spark_examples.py b/test/e2e/spark/test_spark_examples.py index 7a345fb83..561206d87 100644 --- a/test/e2e/spark/test_spark_examples.py +++ b/test/e2e/spark/test_spark_examples.py @@ -111,7 +111,11 @@ def test_spark_connect_crd_smoke(): info = backend._create_session(options=[Name(name)]) assert info.name == name assert info.namespace == namespace - assert info.state in (SparkConnectState.PROVISIONING, SparkConnectState.READY) + assert info.state in ( + SparkConnectState.NEW, + SparkConnectState.PROVISIONING, + SparkConnectState.READY, + ) assert backend.get_session(name).name == name backend.delete_session(name)