Repository navigation
Detect Spark driver completion by container state when tracking via k8s API #68048
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
a4aab1c
6f48fa5
5626137
5577291
5b71389
03c75b9
192d612
5849fe7
ef669f3
4bba0d6
9f88c5b
329d105
2da4998
0e05841
4a7e1f6
24119ab
2b41f9c
80c7d3f
1950310
15ddc22
c70002f
b27c2dc
6b5dc35
e8da750
04004f7
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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. | ||
|
|
||
| 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. | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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) |
||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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" | ||
|
|
@@ -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 {} | ||
|
|
@@ -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 | ||
|
|
@@ -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: | ||
| """ | ||
|
|
@@ -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: | ||
|
|
@@ -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 | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 |
||
| 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) | ||
|
|
@@ -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, | ||
|
|
@@ -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 | ||
|
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." | ||
| ) | ||
|
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": | ||
|
|
@@ -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: | ||
|
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, | ||
| ) | ||
|
|
@@ -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 | ||
|
|
@@ -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: | ||
|
|
@@ -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() | ||
|
|
@@ -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 | ||
|
|
||
There was a problem hiding this comment.
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.