From a4aab1c778148a269c826f80361376641f856770 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Thu, 4 Jun 2026 20:48:25 -0400 Subject: [PATCH 01/15] Use container state to determine spark job status --- .../apache/spark/hooks/spark_submit.py | 94 +++++++++++++------ 1 file changed, 63 insertions(+), 31 deletions(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index 7cf1f3248adc8..f419927c56c73 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1130,7 +1130,7 @@ def _poll_k8s_driver_via_api(self) -> None: consecutive_api_errors = 0 max_consecutive_api_errors = 3 consecutive_pending = 0 - pending_warn_threshold = 10 + waiting_or_pending_warn_threshold = 10 try: if not pod_name: @@ -1162,39 +1162,71 @@ def _poll_k8s_driver_via_api(self) -> None: time.sleep(poll_interval) continue - phase = pod.status.phase or "Initializing" - self.log.info("Application status for %s (phase: %s)", app_id, phase) - if phase == "Succeeded": - break - if phase == "Failed": - container_state = "" - if pod.status.container_statuses: - cs = pod.status.container_statuses[0] - if cs.state and cs.state.terminated: - container_state = f" exit_code={cs.state.terminated.exit_code} reason={cs.state.terminated.reason}" - raise RuntimeError(f"Spark application {app_id} failed (phase=Failed{container_state})") - if phase == "Pending": - consecutive_pending += 1 - if consecutive_pending == pending_warn_threshold: - self.log.warning( - "Driver pod %s has been Pending for %d polls (~%ds); " - "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", - pod_name, - consecutive_pending, - consecutive_pending * poll_interval, - ) + for container in pod.spec.containers: + if "spark" in container.name.lower() or "driver" in container.name.lower(): + driver_container = container + break + if len(pod.spec.containers) == 1: + driver_container = container + else: + driver_container + None + + if driver_container: + for status in pod.status.container_statuses or []: + if status.name == driver_container.name: + if status.state and status.state.terminated: + driver_exit_code = status.state.terminated.exit_code + if driver_exit_code == 0: + break + raise RuntimeError( + f"Spark application {app_id} failed.\nThe driver container exited with a non-zero status code.\nExit code: {driver_exit_code}\nReason: {status.state.terminated.reason}" + ) + if status.state and status.state.waiting: + consecutive_waiting += 1 + if consecutive_waiting == waiting_or_pending_warn_threshold: + self.log.warning( + "Driver container %s has been waiting for %d polls (~%ds); " + "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", + driver_container.name, + consecutive_waiting, + consecutive_waiting * poll_interval, + ) else: - consecutive_pending = 0 - - if phase == "Unknown": - consecutive_unknown += 1 - if consecutive_unknown >= max_consecutive_unknown: + phase = pod.status.phase or "Initializing" + self.log.info("Application status for %s (phase: %s)", app_id, phase) + if phase == "Succeeded": + break + if phase == "Failed": + container_state = "" + if pod.status.container_statuses: + cs = pod.status.container_statuses[0] + if cs.state and cs.state.terminated: + container_state = f" exit_code={cs.state.terminated.exit_code} reason={cs.state.terminated.reason}" raise RuntimeError( - f"Spark application {app_id} reported Unknown phase " - f"{consecutive_unknown} times consecutively; giving up." + f"Spark application {app_id} failed (phase=Failed{container_state})" ) - else: - consecutive_unknown = 0 + if phase == "Pending": + consecutive_pending += 1 + if consecutive_pending == waiting_or_pending_warn_threshold: + self.log.warning( + "Driver pod %s has been Pending for %d polls (~%ds); " + "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", + pod_name, + consecutive_pending, + consecutive_pending * poll_interval, + ) + else: + consecutive_pending = 0 + + if phase == "Unknown": + consecutive_unknown += 1 + if consecutive_unknown >= max_consecutive_unknown: + raise RuntimeError( + f"Spark application {app_id} reported Unknown phase " + f"{consecutive_unknown} times consecutively; giving up." + ) + else: + consecutive_unknown = 0 time.sleep(poll_interval) self._delete_driver_pod() finally: From 6f48fa5e9564c8ecd16ad64205b5c96f401c0c50 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Thu, 4 Jun 2026 20:54:52 -0400 Subject: [PATCH 02/15] Correct typo and conditional logic --- .../src/airflow/providers/apache/spark/hooks/spark_submit.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index f419927c56c73..f427a2f0bef8e 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1169,7 +1169,7 @@ def _poll_k8s_driver_via_api(self) -> None: if len(pod.spec.containers) == 1: driver_container = container else: - driver_container + None + driver_container = None if driver_container: for status in pod.status.container_statuses or []: @@ -1191,6 +1191,8 @@ def _poll_k8s_driver_via_api(self) -> None: consecutive_waiting, consecutive_waiting * poll_interval, ) + else: + consecutive_waiting = 0 else: phase = pod.status.phase or "Initializing" self.log.info("Application status for %s (phase: %s)", app_id, phase) From 5626137b9fe07bc70753ce25ee845ee34941bdb0 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 9 Jun 2026 11:01:21 -0400 Subject: [PATCH 03/15] Correct failing tests --- .../apache/spark/hooks/spark_submit.py | 4 ++-- .../apache/spark/hooks/test_spark_submit.py | 18 +++++++++++------- 2 files changed, 13 insertions(+), 9 deletions(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index f427a2f0bef8e..f6e767d4d45e0 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1130,6 +1130,7 @@ def _poll_k8s_driver_via_api(self) -> None: consecutive_api_errors = 0 max_consecutive_api_errors = 3 consecutive_pending = 0 + consecutive_waiting = 0 waiting_or_pending_warn_threshold = 10 try: @@ -1161,6 +1162,7 @@ def _poll_k8s_driver_via_api(self) -> None: ) from e time.sleep(poll_interval) continue + driver_container = None for container in pod.spec.containers: if "spark" in container.name.lower() or "driver" in container.name.lower(): @@ -1168,8 +1170,6 @@ def _poll_k8s_driver_via_api(self) -> None: break if len(pod.spec.containers) == 1: driver_container = container - else: - driver_container = None if driver_container: for status in pod.status.container_statuses or []: diff --git a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py index f4a610a940814..23fea8ce285de 100644 --- a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py @@ -27,7 +27,7 @@ import kubernetes import pytest import requests -from kubernetes.client import V1Pod, V1PodStatus +from kubernetes.client import V1Pod, V1PodSpec, V1PodStatus from airflow.models import Connection from airflow.providers.apache.spark.hooks.spark_submit import SparkSubmitHook @@ -1503,8 +1503,8 @@ def test_poll_k8s_driver_succeeds(self, mock_get_client): hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value - running_pod = V1Pod(status=V1PodStatus(phase="Running")) - succeeded_pod = V1Pod(status=V1PodStatus(phase="Succeeded")) + running_pod = V1Pod(spec=V1PodSpec(containers=[]), status=V1PodStatus(phase="Running")) + succeeded_pod = V1Pod(spec=V1PodSpec(containers=[]), status=V1PodStatus(phase="Succeeded")) mock_client.read_namespaced_pod.side_effect = [running_pod, succeeded_pod] with patch.object(hook, "_run_post_submit_commands"): @@ -1519,7 +1519,7 @@ def test_poll_k8s_driver_raises_on_failed(self, mock_get_client): hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value - failed_pod = V1Pod(status=V1PodStatus(phase="Failed")) + failed_pod = V1Pod(spec=V1PodSpec(containers=[]), status=V1PodStatus(phase="Failed")) mock_client.read_namespaced_pod.return_value = failed_pod with pytest.raises(RuntimeError, match="phase=Failed"): @@ -1532,7 +1532,9 @@ def test_poll_k8s_driver_raises_after_consecutive_unknown(self, mock_get_client) hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value - mock_client.read_namespaced_pod.return_value = V1Pod(status=V1PodStatus(phase="Unknown")) + mock_client.read_namespaced_pod.return_value = V1Pod( + spec=V1PodSpec(containers=[]), status=V1PodStatus(phase="Unknown") + ) with patch("time.sleep"), pytest.raises(RuntimeError, match="Unknown phase"): hook._poll_k8s_driver_via_api() @@ -1549,7 +1551,7 @@ def test_poll_k8s_driver_tolerates_transient_api_errors(self, mock_get_client, _ mock_client = mock_get_client.return_value api_error = kube_client.ApiException(status=500, reason="Internal Server Error") - succeeded_pod = V1Pod(status=V1PodStatus(phase="Succeeded")) + succeeded_pod = V1Pod(spec=V1PodSpec(containers=[]), status=V1PodStatus(phase="Succeeded")) mock_client.read_namespaced_pod.side_effect = [api_error, api_error, succeeded_pod] with patch.object(hook, "_run_post_submit_commands"): @@ -1565,7 +1567,9 @@ def test_post_submit_commands_run_exactly_once_on_k8s_path(self, mock_get_client hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value - mock_client.read_namespaced_pod.return_value = V1Pod(status=V1PodStatus(phase="Succeeded")) + mock_client.read_namespaced_pod.return_value = V1Pod( + spec=V1PodSpec(containers=[]), status=V1PodStatus(phase="Succeeded") + ) with patch.object(hook, "_run_post_submit_commands") as mock_cmd: hook._poll_k8s_driver_via_api() From 5b713894cbcf35acbbb27018731ead682b6fd70b Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 9 Jun 2026 12:32:32 -0400 Subject: [PATCH 04/15] Add unit tests and fix bugs --- .../apache/spark/hooks/spark_submit.py | 5 +- .../apache/spark/hooks/test_spark_submit.py | 147 +++++++++++++++++- 2 files changed, 149 insertions(+), 3 deletions(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index f6e767d4d45e0..1d08117758a67 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1170,13 +1170,14 @@ def _poll_k8s_driver_via_api(self) -> None: break if len(pod.spec.containers) == 1: driver_container = container - + container_completed = False if driver_container: for status in pod.status.container_statuses or []: if status.name == driver_container.name: if status.state and status.state.terminated: driver_exit_code = status.state.terminated.exit_code if driver_exit_code == 0: + container_completed = True break raise RuntimeError( f"Spark application {app_id} failed.\nThe driver container exited with a non-zero status code.\nExit code: {driver_exit_code}\nReason: {status.state.terminated.reason}" @@ -1229,6 +1230,8 @@ def _poll_k8s_driver_via_api(self) -> None: ) else: consecutive_unknown = 0 + if container_completed: + break time.sleep(poll_interval) self._delete_driver_pod() finally: diff --git a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py index 23fea8ce285de..112b751b5dc6d 100644 --- a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py @@ -22,12 +22,21 @@ from io import StringIO from pathlib import Path from types import ModuleType -from unittest.mock import MagicMock, call, mock_open, patch +from unittest.mock import ANY, MagicMock, call, mock_open, patch import kubernetes import pytest import requests -from kubernetes.client import V1Pod, V1PodSpec, V1PodStatus +from kubernetes.client import ( + V1Container, + V1ContainerState, + V1ContainerStateTerminated, + V1ContainerStateWaiting, + V1ContainerStatus, + V1Pod, + V1PodSpec, + V1PodStatus, +) from airflow.models import Connection from airflow.providers.apache.spark.hooks.spark_submit import SparkSubmitHook @@ -1606,6 +1615,140 @@ def test_poll_k8s_driver_exits_cleanly_on_404(self, mock_get_client): mock_client.delete_namespaced_pod.assert_not_called() + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_container_exit_zero_succeeds(self, mock_get_client): + """Driver container exits cleanly with code 0""" + hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + terminated = V1ContainerStateTerminated(exit_code=0) + state = V1ContainerState(terminated=terminated) + container_status = V1ContainerStatus( + name="spark-driver", + state=state, + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + pod = V1Pod( + spec=V1PodSpec(containers=[V1Container(name="spark-driver")]), + status=V1PodStatus(phase="Running", container_statuses=[container_status]), + ) + mock_client.read_namespaced_pod.return_value = pod + + with patch.object(hook, "_run_post_submit_commands"): + hook._poll_k8s_driver_via_api() + + assert mock_client.read_namespaced_pod.call_count == 1 + + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_container_nonzero_exit_raises(self, mock_get_client): + """Driver container raises RuntimeError with non-zero exit code""" + hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + terminated = V1ContainerStateTerminated(exit_code=1, reason="Error") + state = V1ContainerState(terminated=terminated) + container_status = V1ContainerStatus( + name="spark-driver", + state=state, + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + pod = V1Pod( + spec=V1PodSpec(containers=[V1Container(name="spark-driver")]), + status=V1PodStatus(phase="Running", container_statuses=[container_status]), + ) + mock_client.read_namespaced_pod.return_value = pod + + with pytest.raises(RuntimeError, match="Exit code: 1"): + hook._poll_k8s_driver_via_api() + + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_single_container_fallback(self, mock_get_client): + """Single container with no 'spark' or 'driver' in its name is still set as the driver container""" + hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + terminated = V1ContainerStateTerminated(exit_code=0) + state = V1ContainerState(terminated=terminated) + container_status = V1ContainerStatus( + name="main", + state=state, + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + pod = V1Pod( + spec=V1PodSpec(containers=[V1Container(name="main")]), + status=V1PodStatus(phase="Running", container_statuses=[container_status]), + ) + mock_client.read_namespaced_pod.return_value = pod + + with patch.object(hook, "_run_post_submit_commands"): + hook._poll_k8s_driver_via_api() + + assert mock_client.read_namespaced_pod.call_count == 1 + + @patch("time.sleep") + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_container_waiting_warning(self, mock_get_client, _): + hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + waiting_state = V1ContainerState(waiting=V1ContainerStateWaiting(reason="ContainerCreating")) + waiting_status = V1ContainerStatus( + name="spark-driver", + state=waiting_state, + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + waiting_pod = V1Pod( + spec=V1PodSpec(containers=[V1Container(name="spark-driver")]), + status=V1PodStatus(phase="Running", container_statuses=[waiting_status]), + ) + terminated = V1ContainerStateTerminated(exit_code=0) + done_state = V1ContainerState(terminated=terminated) + done_status = V1ContainerStatus( + name="spark-driver", + state=done_state, + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + done_pod = V1Pod( + spec=V1PodSpec(containers=[V1Container(name="spark-driver")]), + status=V1PodStatus(phase="Running", container_statuses=[done_status]), + ) + mock_client.read_namespaced_pod.side_effect = [waiting_pod] * 10 + [done_pod] + + with patch.object(hook, "_run_post_submit_commands"): + with patch.object(hook.log, "warning") as mock_warning: + hook._poll_k8s_driver_via_api() + + mock_warning.assert_any_call( + "Driver container %s has been waiting for %d polls (~%ds); " + "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", + "spark-driver", + 10, + ANY, + ) + @patch("airflow.providers.apache.spark.hooks.spark_submit.subprocess.run") def test_run_post_submit_commands_runs_only_once(self, mock_run): """Calling _run_post_submit_commands twice must execute commands exactly once.""" From 192d612f0f4f297568d7855a61ee38f2a422c597 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Wed, 17 Jun 2026 00:26:52 -0400 Subject: [PATCH 05/15] Fix failing tests --- .../airflow/providers/apache/spark/hooks/spark_submit.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index 2fcd882be19ad..dbca346d71d83 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1166,6 +1166,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: consecutive_pending = 0 consecutive_waiting = 0 waiting_or_pending_warn_threshold = 10 + terminal_phase: str | None = None try: if not pod_name: @@ -1254,7 +1255,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: raise RuntimeError(f"Spark application {app_id} failed (phase=Failed{container_state})") if phase == "Pending": consecutive_pending += 1 - if consecutive_pending == pending_warn_threshold: + if consecutive_pending == waiting_or_pending_warn_threshold: self.log.warning( "Driver pod %s has been Pending for %d polls (~%ds); " "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", @@ -1299,6 +1300,9 @@ def _poll_k8s_driver_via_api(self) -> str | None: else: consecutive_unknown = 0 if container_completed: + # Driver container exited 0 — the application succeeded even if the + # pod phase still reads "Running" at this poll. + terminal_phase = "Succeeded" break time.sleep(poll_interval) # Pod deletion is best-effort cleanup. If it fails (e.g. already garbage collected or RBAC From ef669f35ad1742b9eba658b1ba49ff07974dc553 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 7 Jul 2026 13:37:10 -0400 Subject: [PATCH 06/15] Remove redundant code --- .../apache/spark/hooks/spark_submit.py | 46 +++++-------------- 1 file changed, 12 insertions(+), 34 deletions(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index dbca346d71d83..c98e8488de017 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1229,44 +1229,22 @@ def _poll_k8s_driver_via_api(self) -> str | None: ) else: consecutive_waiting = 0 - - phase = pod.status.phase or "Initializing" - self.log.info("Application status for %s (phase: %s)", app_id, phase) - if phase == "Succeeded": - if pod.status.container_statuses: - cs = pod.status.container_statuses[0] - if cs.state and cs.state.terminated: - t = cs.state.terminated - self.log.info( - "Container final status: exit_code=%s reason=%s started_at=%s finished_at=%s", - t.exit_code, - t.reason, - t.started_at, - t.finished_at, - ) - terminal_phase = phase - break - if phase == "Failed": - container_state = "" - if pod.status.container_statuses: - cs = pod.status.container_statuses[0] - if cs.state and cs.state.terminated: - container_state = f" exit_code={cs.state.terminated.exit_code} reason={cs.state.terminated.reason}" - raise RuntimeError(f"Spark application {app_id} failed (phase=Failed{container_state})") - if phase == "Pending": - consecutive_pending += 1 - if consecutive_pending == waiting_or_pending_warn_threshold: - self.log.warning( - "Driver pod %s has been Pending for %d polls (~%ds); " - "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", - pod_name, - consecutive_pending, - consecutive_pending * poll_interval, - ) else: phase = pod.status.phase or "Initializing" self.log.info("Application status for %s (phase: %s)", app_id, phase) if phase == "Succeeded": + if pod.status.container_statuses: + cs = pod.status.container_statuses[0] + if cs.state and cs.state.terminated: + t = cs.state.terminated + self.log.info( + "Container final status: exit_code=%s reason=%s started_at=%s finished_at=%s", + t.exit_code, + t.reason, + t.started_at, + t.finished_at, + ) + terminal_phase = phase break if phase == "Failed": container_state = "" From 329d105a34a45e8bc907d30b36ae01da42e76656 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Mon, 20 Jul 2026 20:55:13 -0400 Subject: [PATCH 07/15] Safeguard against unschedulable pods --- .../apache/spark/hooks/spark_submit.py | 89 +++++++++---------- 1 file changed, 43 insertions(+), 46 deletions(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index c98e8488de017..5dfd2d27c3eeb 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1140,7 +1140,7 @@ def _start_driver_status_tracking(self) -> None: def _poll_k8s_driver_via_api(self) -> str | None: """ - Poll the K8s driver pod phase until it reaches a terminal state. + Poll the K8s driver container status or pod phase until it reaches a terminal state. Returns the terminal phase string (e.g. ``"Succeeded"``) on normal completion, or ``None`` if the pod vanished mid-poll (404 — likely deleted by ``on_kill``). @@ -1229,54 +1229,51 @@ def _poll_k8s_driver_via_api(self) -> str | None: ) else: consecutive_waiting = 0 + phase = pod.status.phase or "Initializing" + self.log.info("Application status for %s (phase: %s)", app_id, phase) + if phase == "Succeeded": + if pod.status.container_statuses: + cs = pod.status.container_statuses[0] + if cs.state and cs.state.terminated: + t = cs.state.terminated + self.log.info( + "Container final status: exit_code=%s reason=%s started_at=%s finished_at=%s", + t.exit_code, + t.reason, + t.started_at, + t.finished_at, + ) + terminal_phase = phase + break + if phase == "Failed" and not container_completed: + container_state = "" + if pod.status.container_statuses: + cs = pod.status.container_statuses[0] + if cs.state and cs.state.terminated: + container_state = f" exit_code={cs.state.terminated.exit_code} reason={cs.state.terminated.reason}" + raise RuntimeError(f"Spark application {app_id} failed (phase=Failed{container_state})") + if phase == "Pending": + consecutive_pending += 1 + if consecutive_pending == waiting_or_pending_warn_threshold: + self.log.warning( + "Driver pod %s has been Pending for %d polls (~%ds); " + "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", + pod_name, + consecutive_pending, + consecutive_pending * poll_interval, + ) else: - phase = pod.status.phase or "Initializing" - self.log.info("Application status for %s (phase: %s)", app_id, phase) - if phase == "Succeeded": - if pod.status.container_statuses: - cs = pod.status.container_statuses[0] - if cs.state and cs.state.terminated: - t = cs.state.terminated - self.log.info( - "Container final status: exit_code=%s reason=%s started_at=%s finished_at=%s", - t.exit_code, - t.reason, - t.started_at, - t.finished_at, - ) - terminal_phase = phase - break - if phase == "Failed": - container_state = "" - if pod.status.container_statuses: - cs = pod.status.container_statuses[0] - if cs.state and cs.state.terminated: - container_state = f" exit_code={cs.state.terminated.exit_code} reason={cs.state.terminated.reason}" + consecutive_pending = 0 + + if phase == "Unknown": + consecutive_unknown += 1 + if consecutive_unknown >= max_consecutive_unknown: raise RuntimeError( - f"Spark application {app_id} failed (phase=Failed{container_state})" + f"Spark application {app_id} reported Unknown phase " + f"{consecutive_unknown} times consecutively; giving up." ) - if phase == "Pending": - consecutive_pending += 1 - if consecutive_pending == waiting_or_pending_warn_threshold: - self.log.warning( - "Driver pod %s has been Pending for %d polls (~%ds); " - "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", - pod_name, - consecutive_pending, - consecutive_pending * poll_interval, - ) - else: - consecutive_pending = 0 - - if phase == "Unknown": - consecutive_unknown += 1 - if consecutive_unknown >= max_consecutive_unknown: - raise RuntimeError( - f"Spark application {app_id} reported Unknown phase " - f"{consecutive_unknown} times consecutively; giving up." - ) - else: - consecutive_unknown = 0 + else: + consecutive_unknown = 0 if container_completed: # Driver container exited 0 — the application succeeded even if the # pod phase still reads "Running" at this poll. From 2da4998652aa45d4bf15bc6186e254c21b3ad524 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 25 Aug 2026 12:29:13 -0400 Subject: [PATCH 08/15] Address PR comments --- .../apache/spark/hooks/spark_submit.py | 25 ++++++++++++++++++- .../apache/spark/operators/spark_submit.py | 4 +++ 2 files changed, 28 insertions(+), 1 deletion(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index 5dfd2d27c3eeb..149b1d4e15914 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -139,6 +139,7 @@ class SparkSubmitHook(BaseHook, LoggingMixin): ``keytab`` and ``principal`` configured use ``requests-kerberos`` automatically. Defaults to ``None`` (no auth for non-Kerberos connections). + :param kubernetes_driver_container: Name of the Spark driver container used for identification to track status """ conn_name_attr = "conn_id" @@ -269,6 +270,7 @@ def __init__( track_driver_via_k8s_api: bool = False, yarn_track_via_rm_api: bool = False, yarn_rm_auth: AuthBase | None = None, + kubernetes_driver_container: str | None = None, ) -> None: super().__init__() self._conf = conf or {} @@ -327,6 +329,7 @@ def __init__( # `_track_yarn_application` does not re-fetch the Spark connection # (and re-hit any configured Secrets Backend) on every iteration. self._yarn_rm_base_url: str | None = None + self.kubernetes_driver_container = kubernetes_driver_container def _resolve_should_track_driver_status(self) -> bool: """ @@ -1147,6 +1150,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: Raises ``RuntimeError`` on failure phases or unrecoverable API errors. """ pod_name = self._kubernetes_driver_pod + driver_container_name = self.kubernetes_driver_container namespace = self._connection["namespace"] app_id = self._kubernetes_application_id or pod_name @@ -1200,13 +1204,25 @@ def _poll_k8s_driver_via_api(self) -> str | None: driver_container = None for container in pod.spec.containers: - if "spark" in container.name.lower() or "driver" in container.name.lower(): + if driver_container_name: + if driver_container_name.lower() == container.name.lower(): + driver_container = container + break + continue + if "driver" in container.name.lower(): driver_container = container break + if "spark" in container.name.lower(): + driver_container = container if len(pod.spec.containers) == 1: driver_container = container + if driver_container_name and not driver_container: + raise ValueError( + f"The driver container name provided does not match any of the containers in pod {pod_name}" + ) container_completed = False if driver_container: + self.log.info("%s has been identified as the driver container", driver_container.name) for status in pod.status.container_statuses or []: if status.name == driver_container.name: if status.state and status.state.terminated: @@ -1252,6 +1268,13 @@ def _poll_k8s_driver_via_api(self) -> str | None: if cs.state and cs.state.terminated: container_state = f" exit_code={cs.state.terminated.exit_code} reason={cs.state.terminated.reason}" raise RuntimeError(f"Spark application {app_id} failed (phase=Failed{container_state})") + if phase == "Failed" and container_completed: + self.log.warning( + "Driver pod %s reported phase=Failed, but driver container %s exited 0; " + "treating the Spark application as succeeded.", + pod_name, + driver_container.name, + ) if phase == "Pending": consecutive_pending += 1 if consecutive_pending == waiting_or_pending_warn_threshold: diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py index 8a2129e967300..e2736e0b042f0 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py @@ -149,6 +149,7 @@ class SparkSubmitOperator(ResumableJobMixin, BaseOperator): :param durable: When ``True`` (the default), the external job ID is persisted to task state store before polling begins so that a worker crash and retry reconnects to the existing job instead of submitting a fresh one. Set to ``False`` to always submit a new job on retry. + :param driver_container_name: Name of the Spark driver container used for identification to track status """ # Generic key used across all Spark deployment modes (standalone driver ID, @@ -223,6 +224,7 @@ def __init__( "openlineage", "spark_inject_transport_info", fallback=False ), reconnect_on_retry: bool | None = None, + driver_container_name: str | None = None, **kwargs: Any, ) -> None: if reconnect_on_retry is not None: @@ -272,6 +274,7 @@ def __init__( self._track_driver_via_k8s_api = track_driver_via_k8s_api self._openlineage_inject_parent_job_info = openlineage_inject_parent_job_info self._openlineage_inject_transport_info = openlineage_inject_transport_info + self.driver_container_name = driver_container_name def execute(self, context: Context) -> None: """Call the SparkSubmitHook to run the provided spark job.""" @@ -509,4 +512,5 @@ def _get_hook(self) -> SparkSubmitHook: track_driver_via_k8s_api=self._track_driver_via_k8s_api, yarn_track_via_rm_api=self._yarn_track_via_rm_api, yarn_rm_auth=self._yarn_rm_auth, + kubernetes_driver_container=self.driver_container_name, ) From 24119ab298a85bed997a48637e34413dfb1ad348 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 25 Aug 2026 16:42:50 -0400 Subject: [PATCH 09/15] Add tests and fix failing tests --- .../apache/spark/hooks/spark_submit.py | 2 +- .../apache/spark/hooks/test_spark_submit.py | 134 ++++++++++++++++++ .../spark/operators/test_spark_submit.py | 7 + 3 files changed, 142 insertions(+), 1 deletion(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index d6b9acba1b7fa..ff5a42e905f38 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1305,7 +1305,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: if cs.state and cs.state.terminated: container_state = f" exit_code={cs.state.terminated.exit_code} reason={cs.state.terminated.reason}" raise RuntimeError(f"Spark application {app_id} failed (phase=Failed{container_state})") - if phase == "Failed" and container_completed: + if phase == "Failed" and container_completed and driver_container: self.log.warning( "Driver pod %s reported phase=Failed, but driver container %s exited 0; " "treating the Spark application as succeeded.", diff --git a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py index b007c43d575b8..2dcae86895e11 100644 --- a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py @@ -1824,6 +1824,140 @@ def test_poll_k8s_driver_single_container_fallback(self, mock_get_client): assert mock_client.read_namespaced_pod.call_count == 1 + @patch("time.sleep") + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_prioritizes_driver_over_spark_sidecar(self, mock_get_client, _): + """A sidecar matching 'spark' that exits first must not be mistaken for the driver.""" + hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + containers = [V1Container(name="spark-metrics-sidecar"), V1Container(name="app-driver")] + sidecar_done = V1ContainerStatus( + name="spark-metrics-sidecar", + state=V1ContainerState(terminated=V1ContainerStateTerminated(exit_code=0)), + ready=False, + restart_count=0, + image="busybox", + image_id="", + ) + driver_still_running = V1Pod( + spec=V1PodSpec(containers=containers), + status=V1PodStatus(phase="Running", container_statuses=[sidecar_done]), + ) + driver_done = V1ContainerStatus( + name="app-driver", + state=V1ContainerState(terminated=V1ContainerStateTerminated(exit_code=0)), + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + driver_finished = V1Pod( + spec=V1PodSpec(containers=containers), + status=V1PodStatus(phase="Running", container_statuses=[sidecar_done, driver_done]), + ) + mock_client.read_namespaced_pod.side_effect = [driver_still_running, driver_finished] + + with patch.object(hook, "_run_post_submit_commands"), patch.object(hook.log, "info") as mock_info: + result = hook._poll_k8s_driver_via_api() + + assert result == "Succeeded" + assert mock_client.read_namespaced_pod.call_count == 2 + mock_info.assert_any_call("%s has been identified as the driver container", "app-driver") + + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_custom_container_name_override(self, mock_get_client): + """An explicit kubernetes_driver_container matches by exact name, bypassing the heuristic.""" + hook = SparkSubmitHook( + conn_id="spark_k8s_cluster", + track_driver_via_k8s_api=True, + kubernetes_driver_container="custom-driver", + ) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + terminated_status = V1ContainerStatus( + name="custom-driver", + state=V1ContainerState(terminated=V1ContainerStateTerminated(exit_code=0)), + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + pod = V1Pod( + spec=V1PodSpec( + containers=[V1Container(name="custom-driver"), V1Container(name="spark-exporter")] + ), + status=V1PodStatus(phase="Running", container_statuses=[terminated_status]), + ) + mock_client.read_namespaced_pod.return_value = pod + + with patch.object(hook, "_run_post_submit_commands"), patch.object(hook.log, "info") as mock_info: + result = hook._poll_k8s_driver_via_api() + + assert result == "Succeeded" + mock_info.assert_any_call("%s has been identified as the driver container", "custom-driver") + + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_custom_container_name_no_match_raises(self, mock_get_client): + """An override that matches no container on the pod must raise, not fall back silently.""" + hook = SparkSubmitHook( + conn_id="spark_k8s_cluster", + track_driver_via_k8s_api=True, + kubernetes_driver_container="does-not-exist", + ) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + pod = V1Pod( + spec=V1PodSpec(containers=[V1Container(name="spark-driver")]), + status=V1PodStatus(phase="Running"), + ) + mock_client.read_namespaced_pod.return_value = pod + + with pytest.raises(ValueError, match="does not match any of the containers in pod"): + hook._poll_k8s_driver_via_api() + + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") + def test_poll_k8s_driver_failed_phase_with_completed_container_warns(self, mock_get_client): + """phase=Failed with a driver container that exited 0 must warn and succeed, not raise.""" + hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) + hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._kubernetes_application_id = "spark-abc" + + mock_client = mock_get_client.return_value + terminated_status = V1ContainerStatus( + name="spark-driver", + state=V1ContainerState(terminated=V1ContainerStateTerminated(exit_code=0)), + ready=False, + restart_count=0, + image="spark:3", + image_id="", + ) + pod = V1Pod( + spec=V1PodSpec(containers=[V1Container(name="spark-driver")]), + status=V1PodStatus(phase="Failed", container_statuses=[terminated_status]), + ) + mock_client.read_namespaced_pod.return_value = pod + + with ( + patch.object(hook, "_run_post_submit_commands"), + patch.object(hook.log, "warning") as mock_warning, + ): + result = hook._poll_k8s_driver_via_api() + + assert result == "Succeeded" + mock_warning.assert_any_call( + "Driver pod %s reported phase=Failed, but driver container %s exited 0; " + "treating the Spark application as succeeded.", + "spark-app-abc-driver", + "spark-driver", + ) + @patch("time.sleep") @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_container_waiting_warning(self, mock_get_client, _): diff --git a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py index c55cb10f3ebdd..b7e940ee0303a 100644 --- a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py @@ -927,6 +927,13 @@ def _make_k8s_hook(self): hook._conf = {} return hook + def test_get_hook_passes_driver_container_name(self): + operator = self._make_operator(driver_container_name="custom-driver") + + hook = operator._get_hook() + + assert hook.kubernetes_driver_container == "custom-driver" + def test_execute_calls_submit_then_poll_when_flag_set(self): operator = self._make_operator(track_driver_via_k8s_api=True) hook = self._make_k8s_hook() From 2b41f9cadae66e8a485b529d83ffe9a4c3823342 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 25 Aug 2026 16:59:12 -0400 Subject: [PATCH 10/15] Make driver name field templatable and avoid repetitive logging --- .../src/airflow/providers/apache/spark/hooks/spark_submit.py | 5 ++++- .../airflow/providers/apache/spark/operators/spark_submit.py | 1 + .../tests/unit/apache/spark/operators/test_spark_submit.py | 3 +++ 3 files changed, 8 insertions(+), 1 deletion(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index ff5a42e905f38..7eecb9094ffdf 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1208,6 +1208,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: consecutive_waiting = 0 waiting_or_pending_warn_threshold = 10 terminal_phase: str | None = None + driver_container_logged = False try: if not pod_name: @@ -1259,7 +1260,9 @@ def _poll_k8s_driver_via_api(self) -> str | None: ) container_completed = False if driver_container: - self.log.info("%s has been identified as the driver container", driver_container.name) + if not driver_container_logged: + self.log.info("%s has been identified as the driver container", driver_container.name) + driver_container_logged = True for status in pod.status.container_statuses or []: if status.name == driver_container.name: if status.state and status.state.terminated: diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py index f33df80ac2adc..f3c9c19b44209 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py @@ -370,6 +370,7 @@ class SparkSubmitOperator(ResumableJobMixin, BaseOperator): "env_vars", "post_submit_commands", "properties_file", + "driver_container_name", ) def __init__( diff --git a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py index b7e940ee0303a..e1e0e96c3dda3 100644 --- a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py @@ -934,6 +934,9 @@ def test_get_hook_passes_driver_container_name(self): assert hook.kubernetes_driver_container == "custom-driver" + def test_driver_container_name_is_templatable(self): + assert "driver_container_name" in SparkSubmitOperator.template_fields + def test_execute_calls_submit_then_poll_when_flag_set(self): operator = self._make_operator(track_driver_via_k8s_api=True) hook = self._make_k8s_hook() From 195031020b9e90590c4da4905a0802e99693023a Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Thu, 3 Sep 2026 13:03:35 -0400 Subject: [PATCH 11/15] Change var names --- .../apache/spark/hooks/spark_submit.py | 58 +++++++++---------- .../apache/spark/operators/spark_submit.py | 14 ++--- .../apache/spark/hooks/test_spark_submit.py | 42 +++++++------- .../spark/operators/test_spark_submit.py | 20 +++---- 4 files changed, 66 insertions(+), 68 deletions(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index 7eecb9094ffdf..578b57a46bcf4 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -143,7 +143,7 @@ class SparkSubmitHook(BaseHook, LoggingMixin): ``keytab`` and ``principal`` configured use ``requests-kerberos`` automatically. Defaults to ``None`` (no auth for non-Kerberos connections). - :param kubernetes_driver_container: Name of the Spark driver container used for identification to track status + :param k8s_driver_container_name: Name of the Spark driver container used for identification to track status """ conn_name_attr = "conn_id" @@ -274,7 +274,7 @@ def __init__( track_driver_via_k8s_api: bool = False, yarn_track_via_rm_api: bool = False, yarn_rm_auth: AuthBase | None = None, - kubernetes_driver_container: str | None = None, + k8s_driver_container_name: str | None = None, ) -> None: super().__init__() self._conf = conf or {} @@ -304,7 +304,7 @@ def __init__( self._verbose = verbose self._submit_sp: Any | None = None self._yarn_application_id: str | None = None - self._kubernetes_driver_pod: str | None = None + self._k8s_driver_pod_name: str | None = None self._kubernetes_application_id: str | None = None self.spark_binary = spark_binary self._properties_file = properties_file @@ -336,7 +336,7 @@ def __init__( # `_track_yarn_application` does not re-fetch the Spark connection # (and re-hit any configured Secrets Backend) on every iteration. self._yarn_rm_base_url: str | None = None - self.kubernetes_driver_container = kubernetes_driver_container + self.k8s_driver_container_name = k8s_driver_container_name def _resolve_should_track_driver_status(self) -> bool: """ @@ -860,13 +860,13 @@ def _process_spark_submit_log(self, itr: Iterator[Any]) -> None: # "pod name: -driver" and "submission ID spark:-driver" match_driver_pod = re.search(r"\s*pod name: ((.+?)-([a-z0-9]+)-driver$)", line) if match_driver_pod: - self._kubernetes_driver_pod = match_driver_pod.group(1) - self.log.info("Identified spark driver pod: %s", self._kubernetes_driver_pod) - if not self._kubernetes_driver_pod: + self._k8s_driver_pod_name = match_driver_pod.group(1) + self.log.info("Identified spark driver pod: %s", self._k8s_driver_pod_name) + if not self._k8s_driver_pod_name: match_submission_id = re.search(r"submission ID spark:(.+?-driver)", line) if match_submission_id: - self._kubernetes_driver_pod = match_submission_id.group(1) - self.log.info("Identified spark driver pod: %s", self._kubernetes_driver_pod) + self._k8s_driver_pod_name = match_submission_id.group(1) + self.log.info("Identified spark driver pod: %s", self._k8s_driver_pod_name) match_application_id = re.search(r"\s*spark-app-selector -> (spark-([a-z0-9]+)), ", line) if match_application_id: @@ -1186,10 +1186,10 @@ def _poll_k8s_driver_via_api(self) -> str | None: or ``None`` if the pod vanished mid-poll (404 — likely deleted by ``on_kill``). Raises ``RuntimeError`` on failure phases or unrecoverable API errors. """ - pod_name = self._kubernetes_driver_pod - driver_container_name = self.kubernetes_driver_container + k8s_driver_pod_name = self._k8s_driver_pod_name + k8s_driver_container_name = self.k8s_driver_container_name namespace = self._connection["namespace"] - app_id = self._kubernetes_application_id or pod_name + app_id = self._kubernetes_application_id or k8s_driver_pod_name client = kube_client.get_kube_client() poll_interval = max(self._status_poll_interval, 20) @@ -1211,23 +1211,23 @@ def _poll_k8s_driver_via_api(self) -> str | None: driver_container_logged = False try: - if not pod_name: + if not k8s_driver_pod_name: raise ValueError("K8s driver pod name not set; cannot poll status.") while True: try: - pod = client.read_namespaced_pod(pod_name, namespace) + pod = client.read_namespaced_pod(k8s_driver_pod_name, namespace) consecutive_api_errors = 0 except kube_client.ApiException as e: if e.status == 404: self.log.info( "Driver pod %s not found (404); pod was likely deleted by on_kill. Exiting poll loop.", - pod_name, + k8s_driver_pod_name, ) return None consecutive_api_errors += 1 self.log.warning( "ApiException polling pod %s (%d/%d): %s", - pod_name, + k8s_driver_pod_name, consecutive_api_errors, max_consecutive_api_errors, e, @@ -1242,8 +1242,8 @@ def _poll_k8s_driver_via_api(self) -> str | None: driver_container = None for container in pod.spec.containers: - if driver_container_name: - if driver_container_name.lower() == container.name.lower(): + if k8s_driver_container_name: + if k8s_driver_container_name.lower() == container.name.lower(): driver_container = container break continue @@ -1254,9 +1254,9 @@ def _poll_k8s_driver_via_api(self) -> str | None: driver_container = container if len(pod.spec.containers) == 1: driver_container = container - if driver_container_name and not driver_container: + if k8s_driver_container_name and not driver_container: raise ValueError( - f"The driver container name provided does not match any of the containers in pod {pod_name}" + f"The driver container name provided does not match any of the containers in pod {k8s_driver_pod_name}" ) container_completed = False if driver_container: @@ -1312,7 +1312,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: self.log.warning( "Driver pod %s reported phase=Failed, but driver container %s exited 0; " "treating the Spark application as succeeded.", - pod_name, + k8s_driver_pod_name, driver_container.name, ) if phase == "Pending": @@ -1321,7 +1321,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: self.log.warning( "Driver pod %s has been Pending for %d polls (~%ds); " "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", - pod_name, + k8s_driver_pod_name, consecutive_pending, consecutive_pending * poll_interval, ) @@ -1378,20 +1378,18 @@ def _delete_driver_pod(self) -> None: """Delete the Kubernetes driver pod, logging a warning on failure.""" import kubernetes - self.log.info("Deleting driver pod %s on Kubernetes", self._kubernetes_driver_pod) + self.log.info("Deleting driver pod %s on Kubernetes", self._k8s_driver_pod_name) try: client = kube_client.get_kube_client() client.delete_namespaced_pod( - self._kubernetes_driver_pod, + self._k8s_driver_pod_name, self._connection["namespace"], body=kubernetes.client.V1DeleteOptions(), pretty=True, ) - self.log.info("Deleted driver pod %s", self._kubernetes_driver_pod) + self.log.info("Deleted driver pod %s", self._k8s_driver_pod_name) except kube_client.ApiException: - self.log.exception( - "Exception when attempting to delete driver pod %s", self._kubernetes_driver_pod - ) + self.log.exception("Exception when attempting to delete driver pod %s", self._k8s_driver_pod_name) def on_kill(self) -> None: """Kill Spark submit command.""" @@ -1406,7 +1404,7 @@ def on_kill(self) -> None: "Spark driver %s killed with return code: %s", self._driver_id, driver_kill.wait() ) - if self._should_track_driver_via_k8s_api() and self._kubernetes_driver_pod: + if self._should_track_driver_via_k8s_api() and self._k8s_driver_pod_name: # spark-submit exits early under waitAppCompletion=false, so _submit_sp.poll() is # not None during the poll loop — the deletion block below is skipped on kill. self._delete_driver_pod() @@ -1441,7 +1439,7 @@ def on_kill(self) -> None: ) as yarn_kill: self.log.info("YARN app killed with return code: %s", yarn_kill.wait()) - if self._kubernetes_driver_pod: + if self._k8s_driver_pod_name: self._delete_driver_pod() # Opt-in REST kill path — uses the same RM endpoint as polling, no diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py index f3c9c19b44209..f36d65a5cd9c1 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py @@ -110,7 +110,7 @@ class _KubernetesSparkSubmitBackend(_SparkSubmitDeploymentBackend): def submit_job(self, context: Context) -> str | None: self.hook._conf[_K8S_WAIT_APP_COMPLETION_CONF] = "false" self.hook.submit(self.operator.application) - pod_name = self.hook._kubernetes_driver_pod + pod_name = self.hook._k8s_driver_pod_name namespace = self.hook._connection["namespace"] if not pod_name: raise RuntimeError("spark-submit did not capture a K8s driver pod name") @@ -145,7 +145,7 @@ def is_job_succeeded(self, status: str) -> bool: def poll_until_complete(self, external_id: str, context: Context) -> None: if external_id is not None: _, pod_name = self.operator._parse_k8s_external_id(external_id) - self.hook._kubernetes_driver_pod = pod_name + self.hook._k8s_driver_pod_name = pod_name terminal_phase = self.hook._poll_k8s_driver_via_api() # Cache only when the pod actually reached Succeeded, the 404/vanished path # returns None for cases like: pod deleted by on_kill or garbage collected after failure) @@ -341,7 +341,7 @@ class SparkSubmitOperator(ResumableJobMixin, BaseOperator): Requires Airflow 3.3 or newer; below that, ``durable`` has no effect -- setting it explicitly only emits a warning. :param reconnect_on_retry: deprecated, use ``durable`` instead. - :param driver_container_name: Name of the Spark driver container used for identification to track status + :param k8s_driver_container_name: Name of the Spark driver container used for identification to track status """ # Generic key used across all Spark deployment modes (standalone driver ID, @@ -370,7 +370,7 @@ class SparkSubmitOperator(ResumableJobMixin, BaseOperator): "env_vars", "post_submit_commands", "properties_file", - "driver_container_name", + "k8s_driver_container_name", ) def __init__( @@ -418,7 +418,7 @@ def __init__( ), reconnect_on_retry: bool | None = None, durable: bool | None = None, - driver_container_name: str | None = None, + k8s_driver_container_name: str | None = None, **kwargs: Any, ) -> None: if reconnect_on_retry is not None: @@ -474,7 +474,7 @@ def __init__( self._track_driver_via_k8s_api = track_driver_via_k8s_api self._openlineage_inject_parent_job_info = openlineage_inject_parent_job_info self._openlineage_inject_transport_info = openlineage_inject_transport_info - self.driver_container_name = driver_container_name + self.k8s_driver_container_name = k8s_driver_container_name def execute(self, context: Context) -> None: """Call the SparkSubmitHook to run the provided spark job.""" @@ -607,5 +607,5 @@ def _get_hook(self) -> SparkSubmitHook: track_driver_via_k8s_api=self._track_driver_via_k8s_api, yarn_track_via_rm_api=self._yarn_track_via_rm_api, yarn_rm_auth=self._yarn_rm_auth, - kubernetes_driver_container=self.driver_container_name, + k8s_driver_container_name=self.k8s_driver_container_name, ) diff --git a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py index 2dcae86895e11..1fe0b5bec6609 100644 --- a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py @@ -962,7 +962,7 @@ def test_process_spark_submit_log_k8s(self, pod_name): hook._process_spark_submit_log(log_lines) # Then - assert hook._kubernetes_driver_pod == pod_name + assert hook._k8s_driver_pod_name == pod_name assert hook._kubernetes_application_id == "spark-465b868ada474bda82ccb84ab2747fcd" assert hook._spark_exit_code == 999 @@ -987,7 +987,7 @@ def test_process_spark_submit_log_k8s_submission_id_format(self): hook._process_spark_submit_log(log_lines) - assert hook._kubernetes_driver_pod == "arrow-spark-c8e2e29e73db9c93-driver" + assert hook._k8s_driver_pod_name == "arrow-spark-c8e2e29e73db9c93-driver" def test_process_spark_client_mode_submit_log_k8s(self): # Given @@ -1259,7 +1259,7 @@ def test_on_kill_deletes_pod_when_k8s_api_tracking_and_submit_sp_already_exited( has already exited. """ hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" hook._submit_sp = MagicMock() # spark-submit already exited @@ -1632,7 +1632,7 @@ def test_conf_injection_adds_wait_app_completion(self): @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_succeeds(self, mock_get_client): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1648,7 +1648,7 @@ def test_poll_k8s_driver_succeeds(self, mock_get_client): @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_raises_on_failed(self, mock_get_client): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1661,7 +1661,7 @@ def test_poll_k8s_driver_raises_on_failed(self, mock_get_client): @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_raises_after_consecutive_unknown(self, mock_get_client): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1679,7 +1679,7 @@ def test_poll_k8s_driver_raises_after_consecutive_unknown(self, mock_get_client) @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_tolerates_transient_api_errors(self, mock_get_client, _): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1696,7 +1696,7 @@ def test_poll_k8s_driver_tolerates_transient_api_errors(self, mock_get_client, _ def test_post_submit_commands_run_exactly_once_on_k8s_path(self, mock_get_client): """_run_post_submit_commands must fire exactly once: in _poll_k8s_driver_via_api finally.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1713,7 +1713,7 @@ def test_post_submit_commands_run_exactly_once_on_k8s_path(self, mock_get_client @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_raises_after_consecutive_api_errors(self, mock_get_client, _): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1729,7 +1729,7 @@ def test_poll_k8s_driver_raises_after_consecutive_api_errors(self, mock_get_clie def test_poll_k8s_driver_exits_cleanly_on_404(self, mock_get_client): """404 from read_namespaced_pod means pod was deleted by on_kill — should return cleanly, not raise.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1743,7 +1743,7 @@ def test_poll_k8s_driver_exits_cleanly_on_404(self, mock_get_client): def test_poll_k8s_driver_container_exit_zero_succeeds(self, mock_get_client): """Driver container exits cleanly with code 0""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1772,7 +1772,7 @@ def test_poll_k8s_driver_container_exit_zero_succeeds(self, mock_get_client): def test_poll_k8s_driver_container_nonzero_exit_raises(self, mock_get_client): """Driver container raises RuntimeError with non-zero exit code""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1799,7 +1799,7 @@ def test_poll_k8s_driver_container_nonzero_exit_raises(self, mock_get_client): def test_poll_k8s_driver_single_container_fallback(self, mock_get_client): """Single container with no 'spark' or 'driver' in its name is still set as the driver container""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1829,7 +1829,7 @@ def test_poll_k8s_driver_single_container_fallback(self, mock_get_client): def test_poll_k8s_driver_prioritizes_driver_over_spark_sidecar(self, mock_get_client, _): """A sidecar matching 'spark' that exits first must not be mistaken for the driver.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1869,13 +1869,13 @@ def test_poll_k8s_driver_prioritizes_driver_over_spark_sidecar(self, mock_get_cl @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_custom_container_name_override(self, mock_get_client): - """An explicit kubernetes_driver_container matches by exact name, bypassing the heuristic.""" + """An explicit k8s_driver_container_name matches by exact name, bypassing the heuristic.""" hook = SparkSubmitHook( conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True, - kubernetes_driver_container="custom-driver", + k8s_driver_container_name="custom-driver", ) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1907,9 +1907,9 @@ def test_poll_k8s_driver_custom_container_name_no_match_raises(self, mock_get_cl hook = SparkSubmitHook( conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True, - kubernetes_driver_container="does-not-exist", + k8s_driver_container_name="does-not-exist", ) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1926,7 +1926,7 @@ def test_poll_k8s_driver_custom_container_name_no_match_raises(self, mock_get_cl def test_poll_k8s_driver_failed_phase_with_completed_container_warns(self, mock_get_client): """phase=Failed with a driver container that exited 0 must warn and succeed, not raise.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1962,7 +1962,7 @@ def test_poll_k8s_driver_failed_phase_with_completed_container_warns(self, mock_ @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_container_waiting_warning(self, mock_get_client, _): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._kubernetes_driver_pod = "spark-app-abc-driver" + hook._k8s_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value diff --git a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py index e1e0e96c3dda3..428637cb6cd9d 100644 --- a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py @@ -927,15 +927,15 @@ def _make_k8s_hook(self): hook._conf = {} return hook - def test_get_hook_passes_driver_container_name(self): - operator = self._make_operator(driver_container_name="custom-driver") + def test_get_hook_passes_k8s_driver_container_name(self): + operator = self._make_operator(k8s_driver_container_name="custom-driver") hook = operator._get_hook() - assert hook.kubernetes_driver_container == "custom-driver" + assert hook.k8s_driver_container_name == "custom-driver" - def test_driver_container_name_is_templatable(self): - assert "driver_container_name" in SparkSubmitOperator.template_fields + def test_k8s_driver_container_name_is_templatable(self): + assert "k8s_driver_container_name" in SparkSubmitOperator.template_fields def test_execute_calls_submit_then_poll_when_flag_set(self): operator = self._make_operator(track_driver_via_k8s_api=True) @@ -966,7 +966,7 @@ def test_execute_falls_through_to_plain_submit_when_flag_off(self): def test_k8s_submit_job_returns_encoded_external_id(self): operator = self._make_operator(track_driver_via_k8s_api=True) hook = self._make_k8s_hook() - hook._kubernetes_driver_pod = "spark-abc-driver" + hook._k8s_driver_pod_name = "spark-abc-driver" hook._connection = {"namespace": "mynamespace"} operator._hook = hook @@ -979,7 +979,7 @@ def test_k8s_submit_job_returns_encoded_external_id(self): def test_k8s_submit_job_raises_when_pod_name_missing(self): operator = self._make_operator(track_driver_via_k8s_api=True) hook = self._make_k8s_hook() - hook._kubernetes_driver_pod = None + hook._k8s_driver_pod_name = None hook._connection = {"namespace": "mynamespace"} operator._hook = hook @@ -1060,7 +1060,7 @@ def test_k8s_poll_until_complete_sets_pod_name_and_calls_poll_api(self): operator.poll_until_complete("mynamespace:spark-abc-driver", {}) - assert hook._kubernetes_driver_pod == "spark-abc-driver" + assert hook._k8s_driver_pod_name == "spark-abc-driver" hook._poll_k8s_driver_via_api.assert_called_once() @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="task_state_store requires Airflow 3.3+") @@ -1113,7 +1113,7 @@ def test_k8s_execute_persists_pod_id_when_durable(self): """execute() with durable=True stores the pod ID in task_store before polling.""" operator = self._make_operator(track_driver_via_k8s_api=True, durable=True) hook = self._make_k8s_hook() - hook._kubernetes_driver_pod = "spark-abc-driver" + hook._k8s_driver_pod_name = "spark-abc-driver" hook._connection = {"namespace": "mynamespace"} operator._hook = hook task_store = FakeTaskStateStore() @@ -1136,7 +1136,7 @@ def test_k8s_execute_durable_false_does_not_persist_pod_id(self): """execute() with durable=False does not write spark_job_id to task_store.""" operator = self._make_operator(track_driver_via_k8s_api=True, durable=False) hook = self._make_k8s_hook() - hook._kubernetes_driver_pod = "spark-abc-driver" + hook._k8s_driver_pod_name = "spark-abc-driver" hook._connection = {"namespace": "mynamespace"} operator._hook = hook task_store = FakeTaskStateStore() From 15ddc2228b070f0283b98d6bf205497c6d7cc268 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Thu, 3 Sep 2026 13:25:02 -0400 Subject: [PATCH 12/15] Add docs --- providers/apache/spark/docs/operators.rst | 35 ++++++++++++++++++++--- 1 file changed, 31 insertions(+), 4 deletions(-) diff --git a/providers/apache/spark/docs/operators.rst b/providers/apache/spark/docs/operators.rst index 6da0e87fce66a..fc27021fe02f9 100644 --- a/providers/apache/spark/docs/operators.rst +++ b/providers/apache/spark/docs/operators.rst @@ -264,10 +264,37 @@ Python Kubernetes client rather than holding ``spark-submit`` open for the full persisted to task state before polling begins, so a worker crash and retry reconnects to the existing pod instead of submitting a fresh one. Set ``durable=False`` to always submit a fresh driver on retry. -* Pod completion is detected from ``pod.status.phase``. If your driver pods have sidecar - containers (e.g. Istio injection enabled for the driver namespace), the pod phase may not - advance to ``Succeeded`` until all sidecars exit. In that case the poll loop will wait - indefinitely — set ``execution_timeout`` as a hard bound. + +**Sidecar containers and driver container identification** + +Completion is detected from the driver container's own exit code rather than from +``pod.status.phase`` alone. This matters if your driver pods have sidecar containers: the pod +phase may not advance to ``Succeeded`` until every container exits, but the operator identifies +the driver container specifically and finishes as soon as it exits 0, without waiting on +unrelated sidecars. + +By default the driver container is identified by name, preferring a container with ``driver`` in +its name, then one with ``spark`` in its name, falling back to the pod's only container if there +is just one. If this heuristic doesn't match your setup, set ``k8s_driver_container_name`` to +the exact container name: + +.. code-block:: python + + run_spark = SparkSubmitOperator( + task_id="run_spark", + application="local:///opt/spark/examples/jars/spark-examples.jar", + conn_id="spark_k8s", + deploy_mode="cluster", + track_driver_via_k8s_api=True, + k8s_driver_container_name="spark-kubernetes-driver", + ) + +If ``k8s_driver_container_name`` doesn't match any container on the pod, the task fails +immediately with a ``ValueError`` rather than silently falling back to the heuristic. + +If the pod phase reports ``Failed`` but the driver container itself exited 0 (for example, a +sidecar crashed after the driver finished), the operator logs a warning and still treats the task +as succeeded. YARN ResourceManager API tracking """"""""""""""""""""""""""""""""" From c70002fe96074ecb73a84856177cdb44eb532d67 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Tue, 8 Sep 2026 15:21:54 -0400 Subject: [PATCH 13/15] Replace k8s references with kubernetes --- providers/apache/spark/docs/operators.rst | 6 +- .../apache/spark/hooks/spark_submit.py | 58 ++++++++++--------- .../apache/spark/operators/spark_submit.py | 14 ++--- .../apache/spark/hooks/test_spark_submit.py | 42 +++++++------- .../spark/operators/test_spark_submit.py | 20 +++---- 5 files changed, 71 insertions(+), 69 deletions(-) diff --git a/providers/apache/spark/docs/operators.rst b/providers/apache/spark/docs/operators.rst index fc27021fe02f9..2df8383638a77 100644 --- a/providers/apache/spark/docs/operators.rst +++ b/providers/apache/spark/docs/operators.rst @@ -275,7 +275,7 @@ unrelated sidecars. By default the driver container is identified by name, preferring a container with ``driver`` in its name, then one with ``spark`` in its name, falling back to the pod's only container if there -is just one. If this heuristic doesn't match your setup, set ``k8s_driver_container_name`` to +is just one. If this heuristic doesn't match your setup, set ``kubernetes_driver_container_name`` to the exact container name: .. code-block:: python @@ -286,10 +286,10 @@ the exact container name: conn_id="spark_k8s", deploy_mode="cluster", track_driver_via_k8s_api=True, - k8s_driver_container_name="spark-kubernetes-driver", + kubernetes_driver_container_name="spark-kubernetes-driver", ) -If ``k8s_driver_container_name`` doesn't match any container on the pod, the task fails +If ``kubernetes_driver_container_name`` doesn't match any container on the pod, the task fails immediately with a ``ValueError`` rather than silently falling back to the heuristic. If the pod phase reports ``Failed`` but the driver container itself exited 0 (for example, a diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index 578b57a46bcf4..cb43557d906b2 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -143,7 +143,7 @@ class SparkSubmitHook(BaseHook, LoggingMixin): ``keytab`` and ``principal`` configured use ``requests-kerberos`` automatically. Defaults to ``None`` (no auth for non-Kerberos connections). - :param k8s_driver_container_name: Name of the Spark driver container used for identification to track status + :param kubernetes_driver_container_name: Name of the Spark driver container used for identification to track status """ conn_name_attr = "conn_id" @@ -274,7 +274,7 @@ def __init__( track_driver_via_k8s_api: bool = False, yarn_track_via_rm_api: bool = False, yarn_rm_auth: AuthBase | None = None, - k8s_driver_container_name: str | None = None, + kubernetes_driver_container_name: str | None = None, ) -> None: super().__init__() self._conf = conf or {} @@ -304,7 +304,7 @@ def __init__( self._verbose = verbose self._submit_sp: Any | None = None self._yarn_application_id: str | None = None - self._k8s_driver_pod_name: str | None = None + self._kubernetes_driver_pod_name: str | None = None self._kubernetes_application_id: str | None = None self.spark_binary = spark_binary self._properties_file = properties_file @@ -336,7 +336,7 @@ def __init__( # `_track_yarn_application` does not re-fetch the Spark connection # (and re-hit any configured Secrets Backend) on every iteration. self._yarn_rm_base_url: str | None = None - self.k8s_driver_container_name = k8s_driver_container_name + self.kubernetes_driver_container_name = kubernetes_driver_container_name def _resolve_should_track_driver_status(self) -> bool: """ @@ -860,13 +860,13 @@ def _process_spark_submit_log(self, itr: Iterator[Any]) -> None: # "pod name: -driver" and "submission ID spark:-driver" match_driver_pod = re.search(r"\s*pod name: ((.+?)-([a-z0-9]+)-driver$)", line) if match_driver_pod: - self._k8s_driver_pod_name = match_driver_pod.group(1) - self.log.info("Identified spark driver pod: %s", self._k8s_driver_pod_name) - if not self._k8s_driver_pod_name: + self._kubernetes_driver_pod_name = match_driver_pod.group(1) + self.log.info("Identified spark driver pod: %s", self._kubernetes_driver_pod_name) + if not self._kubernetes_driver_pod_name: match_submission_id = re.search(r"submission ID spark:(.+?-driver)", line) if match_submission_id: - self._k8s_driver_pod_name = match_submission_id.group(1) - self.log.info("Identified spark driver pod: %s", self._k8s_driver_pod_name) + self._kubernetes_driver_pod_name = match_submission_id.group(1) + self.log.info("Identified spark driver pod: %s", self._kubernetes_driver_pod_name) match_application_id = re.search(r"\s*spark-app-selector -> (spark-([a-z0-9]+)), ", line) if match_application_id: @@ -1186,10 +1186,10 @@ def _poll_k8s_driver_via_api(self) -> str | None: or ``None`` if the pod vanished mid-poll (404 — likely deleted by ``on_kill``). Raises ``RuntimeError`` on failure phases or unrecoverable API errors. """ - k8s_driver_pod_name = self._k8s_driver_pod_name - k8s_driver_container_name = self.k8s_driver_container_name + kubernetes_driver_pod_name = self._kubernetes_driver_pod_name + kubernetes_driver_container_name = self.kubernetes_driver_container_name namespace = self._connection["namespace"] - app_id = self._kubernetes_application_id or k8s_driver_pod_name + app_id = self._kubernetes_application_id or kubernetes_driver_pod_name client = kube_client.get_kube_client() poll_interval = max(self._status_poll_interval, 20) @@ -1211,23 +1211,23 @@ def _poll_k8s_driver_via_api(self) -> str | None: driver_container_logged = False try: - if not k8s_driver_pod_name: + if not kubernetes_driver_pod_name: raise ValueError("K8s driver pod name not set; cannot poll status.") while True: try: - pod = client.read_namespaced_pod(k8s_driver_pod_name, namespace) + pod = client.read_namespaced_pod(kubernetes_driver_pod_name, namespace) consecutive_api_errors = 0 except kube_client.ApiException as e: if e.status == 404: self.log.info( "Driver pod %s not found (404); pod was likely deleted by on_kill. Exiting poll loop.", - k8s_driver_pod_name, + kubernetes_driver_pod_name, ) return None consecutive_api_errors += 1 self.log.warning( "ApiException polling pod %s (%d/%d): %s", - k8s_driver_pod_name, + kubernetes_driver_pod_name, consecutive_api_errors, max_consecutive_api_errors, e, @@ -1242,8 +1242,8 @@ def _poll_k8s_driver_via_api(self) -> str | None: driver_container = None for container in pod.spec.containers: - if k8s_driver_container_name: - if k8s_driver_container_name.lower() == container.name.lower(): + if kubernetes_driver_container_name: + if kubernetes_driver_container_name.lower() == container.name.lower(): driver_container = container break continue @@ -1254,9 +1254,9 @@ def _poll_k8s_driver_via_api(self) -> str | None: driver_container = container if len(pod.spec.containers) == 1: driver_container = container - if k8s_driver_container_name and not driver_container: + if kubernetes_driver_container_name and not driver_container: raise ValueError( - f"The driver container name provided does not match any of the containers in pod {k8s_driver_pod_name}" + f"The driver container name provided does not match any of the containers in pod {kubernetes_driver_pod_name}" ) container_completed = False if driver_container: @@ -1312,7 +1312,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: self.log.warning( "Driver pod %s reported phase=Failed, but driver container %s exited 0; " "treating the Spark application as succeeded.", - k8s_driver_pod_name, + kubernetes_driver_pod_name, driver_container.name, ) if phase == "Pending": @@ -1321,7 +1321,7 @@ def _poll_k8s_driver_via_api(self) -> str | None: self.log.warning( "Driver pod %s has been Pending for %d polls (~%ds); " "it may be unschedulable. Continuing to wait — set execution_timeout to bound wait time.", - k8s_driver_pod_name, + kubernetes_driver_pod_name, consecutive_pending, consecutive_pending * poll_interval, ) @@ -1378,18 +1378,20 @@ def _delete_driver_pod(self) -> None: """Delete the Kubernetes driver pod, logging a warning on failure.""" import kubernetes - self.log.info("Deleting driver pod %s on Kubernetes", self._k8s_driver_pod_name) + self.log.info("Deleting driver pod %s on Kubernetes", self._kubernetes_driver_pod_name) try: client = kube_client.get_kube_client() client.delete_namespaced_pod( - self._k8s_driver_pod_name, + self._kubernetes_driver_pod_name, self._connection["namespace"], body=kubernetes.client.V1DeleteOptions(), pretty=True, ) - self.log.info("Deleted driver pod %s", self._k8s_driver_pod_name) + self.log.info("Deleted driver pod %s", self._kubernetes_driver_pod_name) except kube_client.ApiException: - self.log.exception("Exception when attempting to delete driver pod %s", self._k8s_driver_pod_name) + self.log.exception( + "Exception when attempting to delete driver pod %s", self._kubernetes_driver_pod_name + ) def on_kill(self) -> None: """Kill Spark submit command.""" @@ -1404,7 +1406,7 @@ def on_kill(self) -> None: "Spark driver %s killed with return code: %s", self._driver_id, driver_kill.wait() ) - if self._should_track_driver_via_k8s_api() and self._k8s_driver_pod_name: + if self._should_track_driver_via_k8s_api() and self._kubernetes_driver_pod_name: # spark-submit exits early under waitAppCompletion=false, so _submit_sp.poll() is # not None during the poll loop — the deletion block below is skipped on kill. self._delete_driver_pod() @@ -1439,7 +1441,7 @@ def on_kill(self) -> None: ) as yarn_kill: self.log.info("YARN app killed with return code: %s", yarn_kill.wait()) - if self._k8s_driver_pod_name: + if self._kubernetes_driver_pod_name: self._delete_driver_pod() # Opt-in REST kill path — uses the same RM endpoint as polling, no diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py index f36d65a5cd9c1..d1c9002ccb227 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/operators/spark_submit.py @@ -110,7 +110,7 @@ class _KubernetesSparkSubmitBackend(_SparkSubmitDeploymentBackend): def submit_job(self, context: Context) -> str | None: self.hook._conf[_K8S_WAIT_APP_COMPLETION_CONF] = "false" self.hook.submit(self.operator.application) - pod_name = self.hook._k8s_driver_pod_name + pod_name = self.hook._kubernetes_driver_pod_name namespace = self.hook._connection["namespace"] if not pod_name: raise RuntimeError("spark-submit did not capture a K8s driver pod name") @@ -145,7 +145,7 @@ def is_job_succeeded(self, status: str) -> bool: def poll_until_complete(self, external_id: str, context: Context) -> None: if external_id is not None: _, pod_name = self.operator._parse_k8s_external_id(external_id) - self.hook._k8s_driver_pod_name = pod_name + self.hook._kubernetes_driver_pod_name = pod_name terminal_phase = self.hook._poll_k8s_driver_via_api() # Cache only when the pod actually reached Succeeded, the 404/vanished path # returns None for cases like: pod deleted by on_kill or garbage collected after failure) @@ -341,7 +341,7 @@ class SparkSubmitOperator(ResumableJobMixin, BaseOperator): Requires Airflow 3.3 or newer; below that, ``durable`` has no effect -- setting it explicitly only emits a warning. :param reconnect_on_retry: deprecated, use ``durable`` instead. - :param k8s_driver_container_name: Name of the Spark driver container used for identification to track status + :param kubernetes_driver_container_name: Name of the Spark driver container used for identification to track status """ # Generic key used across all Spark deployment modes (standalone driver ID, @@ -370,7 +370,7 @@ class SparkSubmitOperator(ResumableJobMixin, BaseOperator): "env_vars", "post_submit_commands", "properties_file", - "k8s_driver_container_name", + "kubernetes_driver_container_name", ) def __init__( @@ -418,7 +418,7 @@ def __init__( ), reconnect_on_retry: bool | None = None, durable: bool | None = None, - k8s_driver_container_name: str | None = None, + kubernetes_driver_container_name: str | None = None, **kwargs: Any, ) -> None: if reconnect_on_retry is not None: @@ -474,7 +474,7 @@ def __init__( self._track_driver_via_k8s_api = track_driver_via_k8s_api self._openlineage_inject_parent_job_info = openlineage_inject_parent_job_info self._openlineage_inject_transport_info = openlineage_inject_transport_info - self.k8s_driver_container_name = k8s_driver_container_name + self.kubernetes_driver_container_name = kubernetes_driver_container_name def execute(self, context: Context) -> None: """Call the SparkSubmitHook to run the provided spark job.""" @@ -607,5 +607,5 @@ def _get_hook(self) -> SparkSubmitHook: track_driver_via_k8s_api=self._track_driver_via_k8s_api, yarn_track_via_rm_api=self._yarn_track_via_rm_api, yarn_rm_auth=self._yarn_rm_auth, - k8s_driver_container_name=self.k8s_driver_container_name, + kubernetes_driver_container_name=self.kubernetes_driver_container_name, ) diff --git a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py index 1fe0b5bec6609..0ce3ee89c4f66 100644 --- a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py @@ -962,7 +962,7 @@ def test_process_spark_submit_log_k8s(self, pod_name): hook._process_spark_submit_log(log_lines) # Then - assert hook._k8s_driver_pod_name == pod_name + assert hook._kubernetes_driver_pod_name == pod_name assert hook._kubernetes_application_id == "spark-465b868ada474bda82ccb84ab2747fcd" assert hook._spark_exit_code == 999 @@ -987,7 +987,7 @@ def test_process_spark_submit_log_k8s_submission_id_format(self): hook._process_spark_submit_log(log_lines) - assert hook._k8s_driver_pod_name == "arrow-spark-c8e2e29e73db9c93-driver" + assert hook._kubernetes_driver_pod_name == "arrow-spark-c8e2e29e73db9c93-driver" def test_process_spark_client_mode_submit_log_k8s(self): # Given @@ -1259,7 +1259,7 @@ def test_on_kill_deletes_pod_when_k8s_api_tracking_and_submit_sp_already_exited( has already exited. """ hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" hook._submit_sp = MagicMock() # spark-submit already exited @@ -1632,7 +1632,7 @@ def test_conf_injection_adds_wait_app_completion(self): @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_succeeds(self, mock_get_client): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1648,7 +1648,7 @@ def test_poll_k8s_driver_succeeds(self, mock_get_client): @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_raises_on_failed(self, mock_get_client): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1661,7 +1661,7 @@ def test_poll_k8s_driver_raises_on_failed(self, mock_get_client): @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_raises_after_consecutive_unknown(self, mock_get_client): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1679,7 +1679,7 @@ def test_poll_k8s_driver_raises_after_consecutive_unknown(self, mock_get_client) @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_tolerates_transient_api_errors(self, mock_get_client, _): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1696,7 +1696,7 @@ def test_poll_k8s_driver_tolerates_transient_api_errors(self, mock_get_client, _ def test_post_submit_commands_run_exactly_once_on_k8s_path(self, mock_get_client): """_run_post_submit_commands must fire exactly once: in _poll_k8s_driver_via_api finally.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1713,7 +1713,7 @@ def test_post_submit_commands_run_exactly_once_on_k8s_path(self, mock_get_client @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_raises_after_consecutive_api_errors(self, mock_get_client, _): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1729,7 +1729,7 @@ def test_poll_k8s_driver_raises_after_consecutive_api_errors(self, mock_get_clie def test_poll_k8s_driver_exits_cleanly_on_404(self, mock_get_client): """404 from read_namespaced_pod means pod was deleted by on_kill — should return cleanly, not raise.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1743,7 +1743,7 @@ def test_poll_k8s_driver_exits_cleanly_on_404(self, mock_get_client): def test_poll_k8s_driver_container_exit_zero_succeeds(self, mock_get_client): """Driver container exits cleanly with code 0""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1772,7 +1772,7 @@ def test_poll_k8s_driver_container_exit_zero_succeeds(self, mock_get_client): def test_poll_k8s_driver_container_nonzero_exit_raises(self, mock_get_client): """Driver container raises RuntimeError with non-zero exit code""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1799,7 +1799,7 @@ def test_poll_k8s_driver_container_nonzero_exit_raises(self, mock_get_client): def test_poll_k8s_driver_single_container_fallback(self, mock_get_client): """Single container with no 'spark' or 'driver' in its name is still set as the driver container""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1829,7 +1829,7 @@ def test_poll_k8s_driver_single_container_fallback(self, mock_get_client): def test_poll_k8s_driver_prioritizes_driver_over_spark_sidecar(self, mock_get_client, _): """A sidecar matching 'spark' that exits first must not be mistaken for the driver.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1869,13 +1869,13 @@ def test_poll_k8s_driver_prioritizes_driver_over_spark_sidecar(self, mock_get_cl @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_custom_container_name_override(self, mock_get_client): - """An explicit k8s_driver_container_name matches by exact name, bypassing the heuristic.""" + """An explicit kubernetes_driver_container_name matches by exact name, bypassing the heuristic.""" hook = SparkSubmitHook( conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True, - k8s_driver_container_name="custom-driver", + kubernetes_driver_container_name="custom-driver", ) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1907,9 +1907,9 @@ def test_poll_k8s_driver_custom_container_name_no_match_raises(self, mock_get_cl hook = SparkSubmitHook( conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True, - k8s_driver_container_name="does-not-exist", + kubernetes_driver_container_name="does-not-exist", ) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1926,7 +1926,7 @@ def test_poll_k8s_driver_custom_container_name_no_match_raises(self, mock_get_cl def test_poll_k8s_driver_failed_phase_with_completed_container_warns(self, mock_get_client): """phase=Failed with a driver container that exited 0 must warn and succeed, not raise.""" hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value @@ -1962,7 +1962,7 @@ def test_poll_k8s_driver_failed_phase_with_completed_container_warns(self, mock_ @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_container_waiting_warning(self, mock_get_client, _): hook = SparkSubmitHook(conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True) - hook._k8s_driver_pod_name = "spark-app-abc-driver" + hook._kubernetes_driver_pod_name = "spark-app-abc-driver" hook._kubernetes_application_id = "spark-abc" mock_client = mock_get_client.return_value diff --git a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py index 428637cb6cd9d..c98e53adfcdb3 100644 --- a/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/operators/test_spark_submit.py @@ -927,15 +927,15 @@ def _make_k8s_hook(self): hook._conf = {} return hook - def test_get_hook_passes_k8s_driver_container_name(self): - operator = self._make_operator(k8s_driver_container_name="custom-driver") + def test_get_hook_passes_kubernetes_driver_container_name(self): + operator = self._make_operator(kubernetes_driver_container_name="custom-driver") hook = operator._get_hook() - assert hook.k8s_driver_container_name == "custom-driver" + assert hook.kubernetes_driver_container_name == "custom-driver" - def test_k8s_driver_container_name_is_templatable(self): - assert "k8s_driver_container_name" in SparkSubmitOperator.template_fields + def test_kubernetes_driver_container_name_is_templatable(self): + assert "kubernetes_driver_container_name" in SparkSubmitOperator.template_fields def test_execute_calls_submit_then_poll_when_flag_set(self): operator = self._make_operator(track_driver_via_k8s_api=True) @@ -966,7 +966,7 @@ def test_execute_falls_through_to_plain_submit_when_flag_off(self): def test_k8s_submit_job_returns_encoded_external_id(self): operator = self._make_operator(track_driver_via_k8s_api=True) hook = self._make_k8s_hook() - hook._k8s_driver_pod_name = "spark-abc-driver" + hook._kubernetes_driver_pod_name = "spark-abc-driver" hook._connection = {"namespace": "mynamespace"} operator._hook = hook @@ -979,7 +979,7 @@ def test_k8s_submit_job_returns_encoded_external_id(self): def test_k8s_submit_job_raises_when_pod_name_missing(self): operator = self._make_operator(track_driver_via_k8s_api=True) hook = self._make_k8s_hook() - hook._k8s_driver_pod_name = None + hook._kubernetes_driver_pod_name = None hook._connection = {"namespace": "mynamespace"} operator._hook = hook @@ -1060,7 +1060,7 @@ def test_k8s_poll_until_complete_sets_pod_name_and_calls_poll_api(self): operator.poll_until_complete("mynamespace:spark-abc-driver", {}) - assert hook._k8s_driver_pod_name == "spark-abc-driver" + assert hook._kubernetes_driver_pod_name == "spark-abc-driver" hook._poll_k8s_driver_via_api.assert_called_once() @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="task_state_store requires Airflow 3.3+") @@ -1113,7 +1113,7 @@ def test_k8s_execute_persists_pod_id_when_durable(self): """execute() with durable=True stores the pod ID in task_store before polling.""" operator = self._make_operator(track_driver_via_k8s_api=True, durable=True) hook = self._make_k8s_hook() - hook._k8s_driver_pod_name = "spark-abc-driver" + hook._kubernetes_driver_pod_name = "spark-abc-driver" hook._connection = {"namespace": "mynamespace"} operator._hook = hook task_store = FakeTaskStateStore() @@ -1136,7 +1136,7 @@ def test_k8s_execute_durable_false_does_not_persist_pod_id(self): """execute() with durable=False does not write spark_job_id to task_store.""" operator = self._make_operator(track_driver_via_k8s_api=True, durable=False) hook = self._make_k8s_hook() - hook._k8s_driver_pod_name = "spark-abc-driver" + hook._kubernetes_driver_pod_name = "spark-abc-driver" hook._connection = {"namespace": "mynamespace"} operator._hook = hook task_store = FakeTaskStateStore() From b27c2dc018c6e22fcf4b6480b07ecec1273fb113 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Wed, 9 Sep 2026 14:28:34 -0400 Subject: [PATCH 14/15] Change docs to resolve conflict --- providers/apache/spark/docs/operators.rst | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/providers/apache/spark/docs/operators.rst b/providers/apache/spark/docs/operators.rst index 2df8383638a77..cb5a88434ffa2 100644 --- a/providers/apache/spark/docs/operators.rst +++ b/providers/apache/spark/docs/operators.rst @@ -268,12 +268,8 @@ Python Kubernetes client rather than holding ``spark-submit`` open for the full **Sidecar containers and driver container identification** Completion is detected from the driver container's own exit code rather than from -``pod.status.phase`` alone. This matters if your driver pods have sidecar containers: the pod -phase may not advance to ``Succeeded`` until every container exits, but the operator identifies -the driver container specifically and finishes as soon as it exits 0, without waiting on -unrelated sidecars. - -By default the driver container is identified by name, preferring a container with ``driver`` in +``pod.status.phase`` alone as long as the driver container can be identified. By default +the driver container is identified by name, preferring a container with ``driver`` in its name, then one with ``spark`` in its name, falling back to the pod's only container if there is just one. If this heuristic doesn't match your setup, set ``kubernetes_driver_container_name`` to the exact container name: @@ -296,6 +292,11 @@ If the pod phase reports ``Failed`` but the driver container itself exited 0 (fo sidecar crashed after the driver finished), the operator logs a warning and still treats the task as succeeded. +If the driver container could not be identified, ``pod.status.phase`` will be used to track +completion. This matters if your driver pods have sidecar containers: the pod +phase may not advance to ``Succeeded`` until every container exits. To avoid indefinite +waits, set ``execution_timeout`` as a hard bound. + YARN ResourceManager API tracking """"""""""""""""""""""""""""""""" From e8da750ed8a499df255b80f1c8a8cc6efcdb65b1 Mon Sep 17 00:00:00 2001 From: Karen Braganza Date: Thu, 17 Sep 2026 16:42:09 -0400 Subject: [PATCH 15/15] Delete pod if provided driver container name is not found --- .../providers/apache/spark/hooks/spark_submit.py | 6 +++++- .../unit/apache/spark/hooks/test_spark_submit.py | 13 ++++++++++++- 2 files changed, 17 insertions(+), 2 deletions(-) diff --git a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py index cb43557d906b2..38e8c21628d86 100644 --- a/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py +++ b/providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py @@ -1255,8 +1255,12 @@ def _poll_k8s_driver_via_api(self) -> str | None: if len(pod.spec.containers) == 1: driver_container = container if kubernetes_driver_container_name and not driver_container: + with contextlib.suppress(Exception): + self._delete_driver_pod() raise ValueError( - f"The driver container name provided does not match any of the containers in pod {kubernetes_driver_pod_name}" + f"The driver container name provided does not match any of the containers in pod " + f"{kubernetes_driver_pod_name}. Deleted the driver pod; check " + f"kubernetes_driver_container_name for a typo." ) container_completed = False if driver_container: diff --git a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py index 0ce3ee89c4f66..57e9e9d5a27fa 100644 --- a/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py +++ b/providers/apache/spark/tests/unit/apache/spark/hooks/test_spark_submit.py @@ -1903,7 +1903,9 @@ def test_poll_k8s_driver_custom_container_name_override(self, mock_get_client): @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_custom_container_name_no_match_raises(self, mock_get_client): - """An override that matches no container on the pod must raise, not fall back silently.""" + """An override that matches no container on the pod must raise, not fall back silently, + and must delete the now-orphaned driver pod before raising. + """ hook = SparkSubmitHook( conn_id="spark_k8s_cluster", track_driver_via_k8s_api=True, @@ -1922,6 +1924,15 @@ def test_poll_k8s_driver_custom_container_name_no_match_raises(self, mock_get_cl with pytest.raises(ValueError, match="does not match any of the containers in pod"): hook._poll_k8s_driver_via_api() + import kubernetes + + mock_client.delete_namespaced_pod.assert_called_once_with( + "spark-app-abc-driver", + "mynamespace", + body=kubernetes.client.V1DeleteOptions(), + pretty=True, + ) + @patch("airflow.providers.cncf.kubernetes.kube_client.get_kube_client") def test_poll_k8s_driver_failed_phase_with_completed_container_warns(self, mock_get_client): """phase=Failed with a driver container that exited 0 must warn and succeed, not raise."""