Nataneljpwd commented on code in PR #71939:
URL: https://github.com/apache/airflow/pull/71939#discussion_r4148935937


##########
providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py:
##########
@@ -1292,19 +1319,39 @@ def _poll_k8s_driver_via_api(self) -> str | None:
 
     def _build_spark_driver_kill_command(self) -> list[str]:
         """
-        Construct the spark-submit command to kill a driver.
+        Construct the command to kill a driver.
 
         :return: full command to kill a driver
         """
-        # Assume that spark-submit is present in the path to the executing user
-        connection_cmd = [self._connection["spark_binary"]]
+        # spark:// indicates Spark standalone cluster mode
+        if "spark://" in self._connection["master"]:
+            # spark-submit --kill derives its REST URL from the master URL 
(binary RPC
+            # port), which cannot serve REST requests — use the REST API 
directly.
+            urls = [
+                f"{base}/v1/submissions/kill/{self._driver_id}"
+                for base in self._get_standalone_rest_base_urls()
+            ]
+            if len(urls) == 1:
+                connection_cmd = [
+                    "/usr/bin/curl",
+                    "-X",
+                    "DELETE",
+                    urls[0],
+                ]
+            else:
+                # HA: try each master in order
+                curl_cmds = " || ".join(f"/usr/bin/curl --fail -X DELETE 
{shlex.quote(u)}" for u in urls)
+                connection_cmd = ["sh", "-c", curl_cmds]

Review Comment:
   this looks awfully similar to the above for the status tracking with very 
minimal changes, maybe splitting into some kind of helper function will do the 
trick here? the if is fine yet the len and others seems a little weird, also 
can't we just join an array with a length of 1 ans remove any trailing ||? 
isn't it way shorter? that way it might not even require splitting into a 
function as there is too little logic there



##########
providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py:
##########
@@ -676,22 +694,31 @@ def _build_track_driver_status_command(self) -> list[str]:
         """
         curl_max_wait_time = 30
         spark_host = self._connection["master"]
-        if spark_host.endswith(":6066"):
-            spark_host = spark_host.replace("spark://", "http://";)
-            connection_cmd = [
-                "/usr/bin/curl",
-                "--max-time",
-                str(curl_max_wait_time),
-                f"{spark_host}/v1/submissions/status/{self._driver_id}",
-            ]
-            self.log.info(connection_cmd)
-
-            # The driver id so we can poll for its status
+        # spark:// indicates Spark standalone cluster mode
+        if "spark://" in spark_host:
             if not self._driver_id:
                 raise AirflowException(
                     "Invalid status: attempted to poll driver status but no 
driver id is known. Giving up."
                 )
 
+            urls = [
+                f"{base}/v1/submissions/status/{self._driver_id}"
+                for base in self._get_standalone_rest_base_urls()
+            ]
+            if len(urls) == 1:
+                connection_cmd = [
+                    "/usr/bin/curl",
+                    "--max-time",
+                    str(curl_max_wait_time),
+                    urls[0],
+                ]
+            else:
+                # HA: try each master in order (mirrors 
_StandaloneSparkSubmitBackend.get_job_status)
+                curl_cmds = " || ".join(
+                    f"/usr/bin/curl --fail --max-time {curl_max_wait_time} 
{shlex.quote(u)}" for u in urls

Review Comment:
   what if curl is not in /usr/bin? maybe we should dynamically search for it 
with either which or command -v?



##########
providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py:
##########
@@ -676,22 +694,31 @@ def _build_track_driver_status_command(self) -> list[str]:
         """
         curl_max_wait_time = 30
         spark_host = self._connection["master"]
-        if spark_host.endswith(":6066"):
-            spark_host = spark_host.replace("spark://", "http://";)
-            connection_cmd = [
-                "/usr/bin/curl",
-                "--max-time",
-                str(curl_max_wait_time),
-                f"{spark_host}/v1/submissions/status/{self._driver_id}",
-            ]
-            self.log.info(connection_cmd)
-
-            # The driver id so we can poll for its status
+        # spark:// indicates Spark standalone cluster mode
+        if "spark://" in spark_host:
             if not self._driver_id:
                 raise AirflowException(
                     "Invalid status: attempted to poll driver status but no 
driver id is known. Giving up."
                 )
 
+            urls = [
+                f"{base}/v1/submissions/status/{self._driver_id}"
+                for base in self._get_standalone_rest_base_urls()
+            ]
+            if len(urls) == 1:
+                connection_cmd = [
+                    "/usr/bin/curl",
+                    "--max-time",
+                    str(curl_max_wait_time),
+                    urls[0],
+                ]
+            else:
+                # HA: try each master in order (mirrors 
_StandaloneSparkSubmitBackend.get_job_status)
+                curl_cmds = " || ".join(
+                    f"/usr/bin/curl --fail --max-time {curl_max_wait_time} 
{shlex.quote(u)}" for u in urls

Review Comment:
   same as in the above curl



##########
providers/apache/spark/src/airflow/providers/apache/spark/hooks/spark_submit.py:
##########
@@ -676,22 +694,31 @@ def _build_track_driver_status_command(self) -> list[str]:
         """
         curl_max_wait_time = 30
         spark_host = self._connection["master"]
-        if spark_host.endswith(":6066"):
-            spark_host = spark_host.replace("spark://", "http://";)
-            connection_cmd = [
-                "/usr/bin/curl",
-                "--max-time",
-                str(curl_max_wait_time),
-                f"{spark_host}/v1/submissions/status/{self._driver_id}",
-            ]
-            self.log.info(connection_cmd)
-
-            # The driver id so we can poll for its status
+        # spark:// indicates Spark standalone cluster mode
+        if "spark://" in spark_host:
             if not self._driver_id:
                 raise AirflowException(
                     "Invalid status: attempted to poll driver status but no 
driver id is known. Giving up."
                 )
 
+            urls = [
+                f"{base}/v1/submissions/status/{self._driver_id}"
+                for base in self._get_standalone_rest_base_urls()
+            ]
+            if len(urls) == 1:
+                connection_cmd = [
+                    "/usr/bin/curl",
+                    "--max-time",
+                    str(curl_max_wait_time),
+                    urls[0],
+                ]
+            else:
+                # HA: try each master in order (mirrors 
_StandaloneSparkSubmitBackend.get_job_status)
+                curl_cmds = " || ".join(
+                    f"/usr/bin/curl --fail --max-time {curl_max_wait_time} 
{shlex.quote(u)}" for u in urls

Review Comment:
   are you sure it works with the ||? can't stderr or stdout be printed and 
ruin how the logs look like?



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to