Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,9 @@

_K8S_WAIT_APP_COMPLETION_CONF = "spark.kubernetes.submission.waitAppCompletion"

# The JVM's default uncaught-exception handler always prints this exact shape.
_EXCEPTION_START_RE = re.compile(r'Exception in thread "[^"]*"')


class SparkSubmitHook(BaseHook, LoggingMixin):
"""
Expand Down Expand Up @@ -319,9 +322,9 @@ def __init__(
self._driver_id: str | None = None
self._driver_status: str | None = None
self._spark_exit_code: int | None = None
# Last few lines of the spark-submit process's own stdout/stderr, so failure
# exceptions can include the actual root cause instead of just an exit code.
# Rolling tail of spark-submit's own output; widens once _EXCEPTION_START_RE fires.
self._last_submit_log_lines: deque[str] = deque(maxlen=20)
self._exception_anchor_seen: bool = False
self._env: dict[str, Any] | None = None
self._post_submit_commands: list[str] = list(post_submit_commands) if post_submit_commands else []
self._post_submit_commands_done: bool = False
Expand Down Expand Up @@ -885,6 +888,10 @@ def _process_spark_submit_log(self, itr: Iterator[Any]) -> None:
self._driver_id = match_driver_id.group(0)
self.log.info("identified spark driver id: %s", self._driver_id)

if not self._exception_anchor_seen and _EXCEPTION_START_RE.search(line):
# Drop the pre-exception banner noise, keep the whole trace from here on.
self._exception_anchor_seen = True
self._last_submit_log_lines = deque(maxlen=500)
self._last_submit_log_lines.append(line)
self.log.info(line)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -382,7 +382,7 @@ def test_spark_process_runcmd(self, mock_popen, sdk_connection_not_found):
@patch("airflow.providers.apache.spark.hooks.spark_submit.subprocess.Popen")
def test_submit_failure_includes_captured_log_tail(self, mock_popen, sdk_connection_not_found):
mock_popen.return_value.stdout = StringIO(
"Exception in thread main: SparkException: bad jar\nsome other line"
'Exception in thread "main" org.apache.spark.SparkException: bad jar\nsome other line'
)
mock_popen.return_value.stderr = StringIO("")
mock_popen.return_value.wait.return_value = 1
Expand All @@ -391,7 +391,7 @@ def test_submit_failure_includes_captured_log_tail(self, mock_popen, sdk_connect

with pytest.raises(AirflowException, match="Last spark-submit output:") as exc_info:
hook.submit()
assert "Exception in thread main: SparkException: bad jar" in str(exc_info.value)
assert 'Exception in thread "main" org.apache.spark.SparkException: bad jar' in str(exc_info.value)

@pytest.mark.db_test
@patch("airflow.providers.apache.spark.hooks.spark_submit.subprocess.Popen")
Expand Down Expand Up @@ -1013,24 +1013,54 @@ def test_process_spark_submit_log_standalone_cluster(self):

assert hook._driver_id == "driver-20171128111415-0001"

def test_process_spark_submit_log_populates_last_submit_log_lines(self):
def test_process_spark_submit_log_captures_from_exception_marker_onward(self):
"""Lines before the uncaught-exception marker are noise and get discarded."""
hook = SparkSubmitHook(conn_id="spark_standalone_cluster")
log_lines = [
"Running Spark using the REST application submission protocol.",
"17/11/28 11:14:15 INFO RestSubmissionClient: Submitting a request "
"to launch an application in spark://spark-standalone-master:6066",
"WARNING: Using incubator modules: jdk.incubator.vector",
"26/07/27 09:43:44 INFO SparkKubernetesClientFactory: Auto-configuring K8S client",
'Exception in thread "main" io.fabric8.kubernetes.client.KubernetesClientException: boom',
"\tat io.fabric8.kubernetes.client.dsl.internal.OperationSupport.handleCreate(OS.java:340)",
]

hook._process_spark_submit_log(log_lines)

assert list(hook._last_submit_log_lines) == log_lines
assert list(hook._last_submit_log_lines) == [line.strip() for line in log_lines[2:]]

def test_process_spark_submit_log_last_submit_log_lines_truncates_to_maxlen(self):
def test_process_spark_submit_log_without_exception_marker_uses_rolling_tail(self):
"""No 'Exception in thread' anywhere -> falls back to the plain last-20 tail."""
hook = SparkSubmitHook(conn_id="spark_standalone_cluster")
log_lines = [f"line {i}" for i in range(25)]
log_lines = [f"plain output line {i}" for i in range(25)]

hook._process_spark_submit_log(log_lines)

assert list(hook._last_submit_log_lines) == log_lines[-20:]

def test_process_spark_submit_log_exception_message_survives_long_stack_trace(self):
"""A message preceding 30+ stack frames must not roll off the buffer."""
hook = SparkSubmitHook(conn_id="spark_standalone_cluster")
message_line = (
'Exception in thread "main" io.fabric8.kubernetes.client.KubernetesClientException: '
'pods "arrow-spark-driver" is forbidden: exceeded quota: spark-demo-quota'
)
log_lines = [message_line] + [f"\tat some.deep.Frame.method{i}(Frame.java:{i})" for i in range(30)]

hook._process_spark_submit_log(log_lines)

assert next(iter(hook._last_submit_log_lines)) == message_line
assert "exceeded quota" in hook._submit_log_tail

def test_process_spark_submit_log_anchor_buffer_truncates_at_500(self):
"""The widened post-anchor buffer still enforces its own maxlen."""
hook = SparkSubmitHook(conn_id="spark_standalone_cluster")
marker_line = 'Exception in thread "main" java.lang.RuntimeException: boom'
log_lines = [marker_line] + [f"line {i}" for i in range(505)]

hook._process_spark_submit_log(log_lines)

assert len(hook._last_submit_log_lines) == 500
assert list(hook._last_submit_log_lines) == [f"line {i}" for i in range(5, 505)]

def test_process_spark_driver_status_log(self):
# Given
hook = SparkSubmitHook(conn_id="spark_standalone_cluster")
Expand Down