Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
25 commits
Select commit Hold shift + click to select a range
a4aab1c
Use container state to determine spark job status
karenbraganz Jun 5, 2026
6f48fa5
Correct typo and conditional logic
karenbraganz Jun 5, 2026
5626137
Correct failing tests
karenbraganz Jun 9, 2026
5577291
Merge branch 'main' into spark-container-status
karenbraganz Jun 9, 2026
5b71389
Add unit tests and fix bugs
karenbraganz Jun 9, 2026
03c75b9
Merge branch 'main' into spark-container-status
karenbraganz Jun 16, 2026
192d612
Fix failing tests
karenbraganz Jun 17, 2026
5849fe7
Merge branch 'main' into spark-container-status
karenbraganz Jun 17, 2026
ef669f3
Remove redundant code
karenbraganz Jul 7, 2026
4bba0d6
Merge branch 'main' into spark-container-status
karenbraganz Jul 8, 2026
9f88c5b
Merge branch 'main' into spark-container-status
karenbraganz Jul 14, 2026
329d105
Safeguard against unschedulable pods
karenbraganz Jul 21, 2026
2da4998
Address PR comments
karenbraganz Aug 25, 2026
0e05841
Merge remote-tracking branch 'upstream/main' into spark-container-status
karenbraganz Aug 25, 2026
4a7e1f6
Merge branch 'main' into spark-container-status
karenbraganz Aug 25, 2026
24119ab
Add tests and fix failing tests
karenbraganz Aug 25, 2026
2b41f9c
Make driver name field templatable and avoid repetitive logging
karenbraganz Aug 25, 2026
80c7d3f
Merge branch 'main' into spark-container-status
karenbraganz Sep 3, 2026
1950310
Change var names
karenbraganz Sep 3, 2026
15ddc22
Add docs
karenbraganz Sep 3, 2026
c70002f
Replace k8s references with kubernetes
karenbraganz Sep 8, 2026
b27c2dc
Change docs to resolve conflict
karenbraganz Sep 9, 2026
6b5dc35
Merge branch 'main' into spark-container-status
karenbraganz Sep 9, 2026
e8da750
Delete pod if provided driver container name is not found
karenbraganz Sep 17, 2026
04004f7
Merge branch 'main' into spark-container-status
karenbraganz Sep 18, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 36 additions & 4 deletions providers/apache/spark/docs/operators.rst
Original file line number Diff line number Diff line change
Expand Up @@ -284,10 +284,42 @@ before polling begins, and a retry reconnects to that pod instead of submitting
conflicts with the flag and a ``ValueError`` will be raised at task start.
* The Airflow worker must be able to reach the Kubernetes API server and have permission to
read and delete pods in the driver's namespace; otherwise pod tracking and cleanup will fail.
* 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.
* Set ``durable=True`` (the default) to enable crash recovery: the driver pod name is
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.

**Sidecar containers and driver container identification**

Completion is detected from the driver container's own exit code rather than from
``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:

.. 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,
kubernetes_driver_container_name="spark-kubernetes-driver",
)

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.
Comment on lines +312 to +313

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This path also deletes the driver pod, which terminates the running Spark application. That’s a destructive consequence of a typo in a templated field, so it should be stated here rather than only in the exception message.


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.

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.

Set ``durable=False`` to always submit a fresh driver on retry.

@uranusjr uranusjr Oct 10, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this line is stray? (It used to belong to the section above, not the new section)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +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: Name of the Spark driver container used for identification to track status
"""

conn_name_attr = "conn_id"
Expand Down Expand Up @@ -273,6 +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_name: str | None = None,
) -> None:
super().__init__()
self._conf = conf or {}
Expand Down Expand Up @@ -302,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._kubernetes_driver_pod_name: str | None = None
self._kubernetes_application_id: str | None = None
self.spark_binary = spark_binary
self._properties_file = properties_file
Expand Down Expand Up @@ -334,6 +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_name = kubernetes_driver_container_name

def _resolve_should_track_driver_status(self) -> bool:
"""
Expand Down Expand Up @@ -857,13 +860,13 @@ def _process_spark_submit_log(self, itr: Iterator[Any]) -> None:
# "pod name: <name>-driver" and "submission ID spark:<name>-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._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._kubernetes_driver_pod = match_submission_id.group(1)
self.log.info("Identified spark driver pod: %s", self._kubernetes_driver_pod)
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:
Expand Down Expand Up @@ -1177,15 +1180,16 @@ 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``).
Raises ``RuntimeError`` on failure phases or unrecoverable API errors.
"""
pod_name = self._kubernetes_driver_pod
kubernetes_driver_pod_name = self._kubernetes_driver_pod_name
kubernetes_driver_container_name = self.kubernetes_driver_container_name

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is a template field, so this can arrive as whitespace from a Jinja render (not uncommon due to how people indent templates). Maybe do a strip(), or only act against None (instead of falsy) instead below?

namespace = self._connection["namespace"]
app_id = self._kubernetes_application_id or 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)
Expand All @@ -1201,26 +1205,29 @@ 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
driver_container_logged = False

try:
if not 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(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.",
pod_name,
kubernetes_driver_pod_name,
)
return None
consecutive_api_errors += 1
self.log.warning(
"ApiException polling pod %s (%d/%d): %s",
pod_name,
kubernetes_driver_pod_name,
consecutive_api_errors,
max_consecutive_api_errors,
e,
Expand All @@ -1232,7 +1239,56 @@ 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 kubernetes_driver_container_name:
if kubernetes_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
Comment thread
karenbraganz marked this conversation as resolved.
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 "
f"{kubernetes_driver_pod_name}. Deleted the driver pod; check "
f"kubernetes_driver_container_name for a typo."
)
Comment thread
karenbraganz marked this conversation as resolved.
container_completed = False
if driver_container:
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:
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 @@ -1249,20 +1305,27 @@ def _poll_k8s_driver_via_api(self) -> str | None:
)
terminal_phase = phase
break
if phase == "Failed":
if phase == "Failed" and not container_completed:
Comment thread
karenbraganz marked this conversation as resolved.
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 == "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.",
kubernetes_driver_pod_name,
driver_container.name,
)
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.",
pod_name,
kubernetes_driver_pod_name,
consecutive_pending,
consecutive_pending * poll_interval,
)
Expand All @@ -1280,6 +1343,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 Expand Up @@ -1314,19 +1382,19 @@ 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._kubernetes_driver_pod_name)
try:
client = kube_client.get_kube_client()
client.delete_namespaced_pod(
self._kubernetes_driver_pod,
self._kubernetes_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._kubernetes_driver_pod_name)
except kube_client.ApiException:
self.log.exception(
"Exception when attempting to delete driver pod %s", self._kubernetes_driver_pod
"Exception when attempting to delete driver pod %s", self._kubernetes_driver_pod_name
)

def on_kill(self) -> None:
Expand All @@ -1342,7 +1410,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._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()
Expand Down Expand Up @@ -1377,7 +1445,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._kubernetes_driver_pod_name:
self._delete_driver_pod()

# Opt-in REST kill path — uses the same RM endpoint as polling, no
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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._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")
Expand Down Expand Up @@ -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._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)
Expand Down Expand Up @@ -341,6 +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 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,
Expand Down Expand Up @@ -369,6 +370,7 @@ class SparkSubmitOperator(ResumableJobMixin, BaseOperator):
"env_vars",
"post_submit_commands",
"properties_file",
"kubernetes_driver_container_name",
)

def __init__(
Expand Down Expand Up @@ -416,6 +418,7 @@ def __init__(
),
reconnect_on_retry: bool | None = None,
durable: bool | None = None,
kubernetes_driver_container_name: str | None = None,
**kwargs: Any,
) -> None:
if reconnect_on_retry is not None:
Expand Down Expand Up @@ -471,6 +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.kubernetes_driver_container_name = kubernetes_driver_container_name

def execute(self, context: Context) -> None:
"""Call the SparkSubmitHook to run the provided spark job."""
Expand Down Expand Up @@ -603,4 +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_name=self.kubernetes_driver_container_name,
)
Loading