Skip to content
Open
Original file line number Diff line number Diff line change
Expand Up @@ -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``).
Expand All @@ -1164,7 +1164,9 @@ def _poll_k8s_driver_via_api(self) -> str | None:
consecutive_api_errors = 0
max_consecutive_api_errors = 3
consecutive_pending = 0
pending_warn_threshold = 10
consecutive_waiting = 0
waiting_or_pending_warn_threshold = 10
terminal_phase: str | None = None

try:
if not pod_name:
Expand Down Expand Up @@ -1195,7 +1197,38 @@ def _poll_k8s_driver_via_api(self) -> str | 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():
driver_container = container
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}"
)
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_waiting = 0
phase = pod.status.phase or "Initializing"
self.log.info("Application status for %s (phase: %s)", app_id, phase)
if phase == "Succeeded":
Expand All @@ -1212,7 +1245,7 @@ def _poll_k8s_driver_via_api(self) -> str | None:
)
terminal_phase = phase
break
if phase == "Failed":
if phase == "Failed" and not container_completed:
container_state = ""
if pod.status.container_statuses:
cs = pod.status.container_statuses[0]
Expand All @@ -1221,7 +1254,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.",
Expand All @@ -1241,6 +1274,11 @@ 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
# denied), suppress the error so terminal_phase is still returned and the task
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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, 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
Expand Down Expand Up @@ -1532,8 +1541,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"):
Expand All @@ -1548,7 +1557,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"):
Expand All @@ -1561,7 +1570,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()
Expand All @@ -1578,7 +1589,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"):
Expand All @@ -1594,7 +1605,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()
Expand Down Expand Up @@ -1631,6 +1644,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."""
Expand Down