Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 27 additions & 4 deletions kubeflow/spark/backends/kubernetes/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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,
Expand Down
95 changes: 95 additions & 0 deletions kubeflow/spark/backends/kubernetes/backend_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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",
[
Expand Down
8 changes: 4 additions & 4 deletions kubeflow/spark/backends/kubernetes/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
54 changes: 46 additions & 8 deletions kubeflow/spark/backends/kubernetes/utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
8 changes: 6 additions & 2 deletions kubeflow/spark/types/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
4 changes: 2 additions & 2 deletions kubeflow/spark/types/types_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
],
Expand All @@ -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,
],
Expand Down
6 changes: 5 additions & 1 deletion test/e2e/spark/test_spark_examples.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down