This is an automated email from the ASF dual-hosted git repository.
kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 246e3b231c9 Add `include_traceback` to model-backed retry policies
(#74308)
246e3b231c9 is described below
commit 246e3b231c988bd3cc7466350e8e0f7eb3f75b9d
Author: Kaxil Naik <[email protected]>
AuthorDate: Tue Oct 6 08:29:45 2026 +0100
Add `include_traceback` to model-backed retry policies (#74308)
LLMRetryPolicy and ClassifierRetryPolicy send the model only
"ExceptionType: message", so the exception chain, the stack, and
module-qualified class names never reach it. A JSONDecodeError raised while
handling a truncated read looks like bad data, though the chained transport
error makes it a retryable network failure.
include_traceback=True sends traceback.format_exception() output instead,
redacted as a whole and then truncated to max_exception_length keeping the
tail, where the innermost frames and the final exception line are. The default
is unchanged.
---
providers/common/ai/docs/retry_policies.rst | 99 +++++++++++++++---
.../airflow/providers/common/ai/policies/retry.py | 56 +++++++---
.../ai/tests/unit/common/ai/policies/test_retry.py | 116 ++++++++++++++++++++-
3 files changed, 245 insertions(+), 26 deletions(-)
diff --git a/providers/common/ai/docs/retry_policies.rst
b/providers/common/ai/docs/retry_policies.rst
index 79625373783..686eb5c7e0f 100644
--- a/providers/common/ai/docs/retry_policies.rst
+++ b/providers/common/ai/docs/retry_policies.rst
@@ -113,10 +113,11 @@ How it works
When a task fails, either policy:
-1. Sends the exception message to the configured LLM. By default, the message
- is first masked through Airflow's secrets masker (see ``redactor`` below)
- and truncated to ``max_exception_length`` characters before it is added
- to the prompt.
+1. Sends the exception's class name and message to the configured LLM, or the
+ formatted traceback with ``include_traceback=True`` (see
+ `Sending the traceback`_). By default, the text is first masked through
+ Airflow's secrets masker (see ``redactor`` below) and truncated to
+ ``max_exception_length`` characters before it is added to the prompt.
2. With ``LLMRetryPolicy``, the model returns an
:class:`~airflow.providers.common.ai.policies.retry.ErrorClassification`: a
category, whether to retry, a suggested delay, and its reasoning. With
@@ -397,8 +398,9 @@ What the model can and cannot do
Under either policy the model is given no tools and there is no way to attach
any, so it cannot run code, call an API, read a connection, or reach your data.
-Beyond your ``instructions``, it sees only the exception's class name, the
-exception message (after redaction and truncation), how many attempts are left,
+Beyond your ``instructions``, it sees only the exception's class name and
+message (or, with ``include_traceback=True``, the formatted traceback; either
+after redaction and truncation), how many attempts are left,
and, under ``ClassifierRetryPolicy``, the category names and descriptions. The
prompt says
``attempt {try_number} of {max_tries}``, so the model knows the limit and not
just where it is right now; an instruction like "after two attempts treat an
@@ -567,8 +569,9 @@ Both policies share every parameter below except
``categories``,
apply.
* - ``redactor``
- None (uses ``redact_registered_secrets``)
- - Callable ``(str) -> str`` applied to the exception's string
- representation before it is added to the classification prompt. The
+ - Callable ``(str) -> str`` applied to the exception text (its string
+ representation, or the whole traceback with ``include_traceback=True``)
+ before it is added to the classification prompt. The
default only masks values already registered via ``mask_secret()``
(e.g. connection passwords Airflow captured while resolving the
failing task's connections) -- it is not general-purpose PII
@@ -577,15 +580,87 @@ Both policies share every parameter below except
``categories``,
the default masker entirely rather than stacking on top of it.
* - ``redact_exception``
- True
- - Whether to redact the exception's string representation before it is
+ - Whether to redact the exception text before it is
added to the classification prompt. Set to ``False`` to disable
redaction entirely. Raises ``ValueError`` at construction time if
combined with an explicit ``redactor``.
* - ``max_exception_length``
- 4096
- Maximum number of characters of the (already redacted) exception
- message included in the prompt. Longer messages are truncated with a
- trailing ``"... (truncated)"`` marker. Must be a positive integer.
+ text included in the prompt. A longer message keeps its head, with a
+ trailing ``"... (truncated)"`` marker; a longer traceback keeps its
+ tail, with a leading ``"(truncated) ..."`` marker. Must be a positive
+ integer.
+ * - ``include_traceback``
+ - False
+ - Send the formatted traceback, with chained exceptions and
+ module-qualified class names, instead of ``ExceptionType: message``.
+ See `Sending the traceback`_.
+
+Sending the traceback
+---------------------
+
+By default the model sees ``ExceptionType: message``, and some failures name
the
+wrong cause there. A response cut off mid-body and then parsed fails as:
+
+.. code-block:: text
+
+ JSONDecodeError: Expecting ',' delimiter: line 1 column 51 (char 50)
+
+That reads as bad input data, a ``data`` failure that is not retried. The real
+cause is in the exception chain, which the message does not carry. With
+``include_traceback=True`` the policy sends the formatted traceback instead:
+
+.. code-block:: python
+
+ import json
+ from http.client import IncompleteRead
+ from urllib.request import urlopen
+
+ from airflow.providers.common.ai.policies.retry import LLMRetryPolicy
+ from airflow.sdk import task
+
+
+ @task(retries=3,
retry_policy=LLMRetryPolicy(llm_conn_id="pydanticai_default",
include_traceback=True))
+ def fetch_orders():
+ with urlopen("https://api.example.com/orders") as response:
+ try:
+ body = response.read()
+ except IncompleteRead as err:
+ return json.loads(err.partial) # salvage whatever arrived
+ return json.loads(body)
+
+The model then receives both exceptions, with module-qualified class names and
+the linking line between them (stack frames shortened to ``...`` here):
+
+.. code-block:: text
+
+ Traceback (most recent call last):
+ ...
+ http.client.IncompleteRead: IncompleteRead(50 bytes read, 4096 more
expected)
+
+ During handling of the above exception, another exception occurred:
+
+ Traceback (most recent call last):
+ ...
+ json.decoder.JSONDecodeError: Expecting ',' delimiter: line 1 column 51
(char 50)
+
+The ``IncompleteRead`` underneath says the connection dropped mid-body, a
+``network`` failure worth retrying.
+
+The text is what :func:`traceback.format_exception` produces: each frame's file
+path and source line, and the message of every chained exception
+(``raise ... from ...`` and an exception raised while handling another). Local
+variable values are not included. ``redactor`` runs over the whole text before
+it is truncated, so a secret in a chained exception's message is masked like
+one in the final message.
+
+A traceback is often many times longer than the message, and every failure pays
+for it, up to ``max_exception_length`` characters per classification. When it
is
+longer than that, the policy keeps the tail behind a leading ``(truncated)
...``
+marker, because the innermost frames and the final exception line say the most.
+A chained cause is printed first, so a deep stack can push it out of the
+window; raise ``max_exception_length`` if the causes you need are being cut.
Custom redactors
----------------
@@ -615,7 +690,7 @@ yourself if you still want known-secret masking too:
llm_policy = LLMRetryPolicy(
llm_conn_id="pydanticai_default",
redactor=redact_emails_and_secrets,
- max_exception_length=2048, # keep long tracebacks from inflating
token cost
+ max_exception_length=2048, # keep long exception messages from
inflating token cost
)
To disable redaction entirely (for example, if you are certain your
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py
b/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py
index 6440dcbaec6..cccdfbe4593 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py
@@ -37,6 +37,7 @@ Requires Airflow 3.3+ (RetryPolicy was added in AIP-105).
from __future__ import annotations
import logging
+import traceback
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import timedelta
@@ -223,8 +224,10 @@ def redact_registered_secrets(message: str) -> str:
_REDACTION_PARAMS_DOC = """
- :param redactor: Callable applied to the exception's string representation
- before it is added to the classification prompt. Defaults to
+ :param redactor: Callable applied to the exception text (its string
+ representation, or the whole formatted traceback with
+ ``include_traceback=True``) before it is added to the classification
+ prompt. Defaults to
:func:`~airflow.providers.common.ai.policies.retry.redact_registered_secrets`,
which only masks values already registered via ``mask_secret()``.
Pass a custom callable to replace the default masking entirely --
@@ -238,17 +241,31 @@ _REDACTION_PARAMS_DOC = """
an explicit ``redactor`` raises ``ValueError`` at construction time,
since the two settings would otherwise conflict silently.
:param max_exception_length: Maximum number of characters of the
- (already redacted) exception message included in the prompt. Longer
- messages are truncated with a trailing ``"... (truncated)"`` marker.
- Must be a positive integer. Defaults to 4096.
+ (already redacted) exception text included in the prompt. A longer
+ message is cut to its head with a trailing ``"... (truncated)"``
marker;
+ a longer traceback (``include_traceback=True``) is cut to its tail
with a
+ leading ``"(truncated) ..."`` marker, so the innermost frames and the
+ final exception line survive. Must be a positive integer. Defaults to
4096.
+ :param include_traceback: Send the formatted traceback instead of
+ ``ExceptionType: message``. Defaults to ``False``. The traceback is
what
+ :func:`traceback.format_exception` produces: the stack frames with
their
+ file paths and source lines, every chained exception (``__cause__`` and
+ ``__context__``), and module-qualified class names such as
+ ``botocore.exceptions.ClientError``. Local variable values of the
frames
+ are not included. The whole text goes through ``redactor`` before it is
+ truncated. A traceback is usually many times longer than the message,
so
+ each classification costs more input tokens, up to
+ ``max_exception_length`` characters.
.. warning::
The exception's string representation is sent to the configured
external LLM provider (OpenAI, Anthropic, Bedrock, Vertex, Ollama,
etc.) as part of the classification prompt, so it may leak whatever
the failing task put in the exception message — connection strings,
- credential fragments, PII, or other secrets. By default the message
- is run through
+ credential fragments, PII, or other secrets. With
+ ``include_traceback=True`` that also covers the messages of chained
+ exceptions, file paths on the worker, and the source line of each
+ frame. By default the text is run through
:func:`~airflow.providers.common.ai.policies.retry.redact_registered_secrets`
via ``redactor``, which masks values already registered via
``mask_secret()`` (for example, connection passwords Airflow
@@ -279,6 +296,7 @@ class _ModelRetryPolicy(RetryPolicy):
redactor: Callable[[str], str] | None = None,
redact_exception: bool = True,
max_exception_length: int = 4096,
+ include_traceback: bool = False,
) -> None:
if max_exception_length <= 0:
raise ValueError(f"max_exception_length must be a positive
integer, got {max_exception_length}")
@@ -299,6 +317,7 @@ class _ModelRetryPolicy(RetryPolicy):
)
self.redact_exception = redact_exception
self.max_exception_length = max_exception_length
+ self.include_traceback = include_traceback
def _hook(self) -> PydanticAIHook:
from airflow.providers.common.ai.hooks.pydantic_ai import
PydanticAIHook
@@ -306,14 +325,23 @@ class _ModelRetryPolicy(RetryPolicy):
return PydanticAIHook(llm_conn_id=self.llm_conn_id,
model_id=self.model_id)
def _prompt(self, exception: BaseException, try_number: int, max_tries:
int) -> str:
+ if self.include_traceback:
+ text = "".join(traceback.format_exception(exception)).rstrip("\n")
+ else:
+ text = str(exception)
# Redact before truncating -- truncating first could cut a registered
secret in half.
- message = self.redactor(str(exception)) if self.redactor is not None
else str(exception)
- if len(message) > self.max_exception_length:
- message = f"{message[: self.max_exception_length]}... (truncated)"
+ if self.redactor is not None:
+ text = self.redactor(text)
+ if len(text) > self.max_exception_length:
+ if self.include_traceback:
+ # Keep the tail: the innermost frames and the final exception
line say the most.
+ text = f"(truncated) ...{text[-self.max_exception_length :]}"
+ else:
+ text = f"{text[: self.max_exception_length]}... (truncated)"
+ if not self.include_traceback:
+ text = f"{type(exception).__name__}: {text}"
return (
- f"Classify this error from a data pipeline task "
- f"(attempt {try_number} of {max_tries}):\n\n"
- f"{type(exception).__name__}: {message}"
+ f"Classify this error from a data pipeline task (attempt
{try_number} of {max_tries}):\n\n{text}"
)
def _run(
@@ -498,6 +526,7 @@ class ClassifierRetryPolicy(_ModelRetryPolicy):
redactor: Callable[[str], str] | None = None,
redact_exception: bool = True,
max_exception_length: int = 4096,
+ include_traceback: bool = False,
) -> None:
super().__init__(
llm_conn_id,
@@ -508,6 +537,7 @@ class ClassifierRetryPolicy(_ModelRetryPolicy):
redactor=redactor,
redact_exception=redact_exception,
max_exception_length=max_exception_length,
+ include_traceback=include_traceback,
)
self.min_confidence = None if min_confidence is None else
check_bar(min_confidence, "min_confidence")
self.categories: dict[str, ErrorCategory] = self._validate_categories(
diff --git a/providers/common/ai/tests/unit/common/ai/policies/test_retry.py
b/providers/common/ai/tests/unit/common/ai/policies/test_retry.py
index 30ca9d2a679..bd34472b8de 100644
--- a/providers/common/ai/tests/unit/common/ai/policies/test_retry.py
+++ b/providers/common/ai/tests/unit/common/ai/policies/test_retry.py
@@ -17,11 +17,13 @@
from __future__ import annotations
import copy
+import json
import logging
import math
+import traceback
import warnings
from datetime import timedelta
-from unittest.mock import MagicMock, patch
+from unittest.mock import MagicMock, create_autospec, patch
import pytest
from pydantic import TypeAdapter, ValidationError
@@ -654,6 +656,47 @@ class TestConfidenceGate:
assert decision.reason == "category=auth confidence=0.05 threshold=n/a
action=fail"
+PROMPT_HEADER = "Classify this error from a data pipeline task (attempt 2 of
4):\n\n"
+
+POLICY_CLASSES = [
+ pytest.param(ClassifierRetryPolicy, id="classifier"),
+ pytest.param(LLMRetryPolicy, id="llm"),
+]
+
+
+def _chained(outer: Exception, cause: Exception, *, explicit: bool = True) ->
Exception:
+ """Raise ``outer`` while handling ``cause`` so both carry a traceback, and
return ``outer``."""
+ try:
+ try:
+ raise cause
+ except type(cause):
+ if explicit:
+ raise outer from cause
+ raise outer
+ except type(outer) as exc:
+ return exc
+
+
+def _truncated_read(*, explicit: bool = True) -> Exception:
+ """A JSON parse failure raised while handling the transport error that cut
the response body short."""
+ body = '{"rows": [{"id": 1}, {"id'
+ return _chained(
+ json.JSONDecodeError("Unterminated string starting at", body, 22),
+ ConnectionResetError("peer closed connection after 25 of 4096 bytes"),
+ explicit=explicit,
+ )
+
+
+def _prompt_for(mock_hook_cls, policy_cls, exception, **kwargs) -> str:
+ """Evaluate ``exception`` under a ``policy_cls`` built with ``kwargs`` and
return the prompt the model got."""
+ if policy_cls is ClassifierRetryPolicy:
+ agent = _install(mock_hook_cls, _agent("data"))
+ else:
+ agent = _install(mock_hook_cls, _open_agent("data",
should_retry=False))
+ policy_cls(llm_conn_id="test", **kwargs).evaluate(exception, try_number=2,
max_tries=4)
+ return agent.run_sync.call_args.args[0]
+
+
class TestPrompt:
@patch(HOOK, autospec=True)
def test_prompt_includes_exception_type_and_message(self, mock_hook_cls):
@@ -767,6 +810,77 @@ class TestPrompt:
assert agent.run_sync.call_args.args[0].endswith("ConnectionError: pw
*** tail")
+ @pytest.mark.parametrize("policy_cls", POLICY_CLASSES)
+ @patch(HOOK, autospec=True)
+ def test_traceback_is_off_by_default(self, mock_hook_cls, policy_cls):
+ error = _truncated_read()
+
+ prompt = _prompt_for(mock_hook_cls, policy_cls, error)
+
+ assert policy_cls(llm_conn_id="test").include_traceback is False
+ assert prompt == f"{PROMPT_HEADER}JSONDecodeError: {error}"
+
+ @pytest.mark.parametrize("policy_cls", POLICY_CLASSES)
+ @pytest.mark.parametrize(
+ ("explicit", "link"),
+ [
+ pytest.param(True, "The above exception was the direct cause",
id="cause"),
+ pytest.param(False, "During handling of the above exception",
id="context"),
+ ],
+ )
+ @patch(HOOK, autospec=True)
+ def test_traceback_carries_the_chain_and_qualified_names(self,
mock_hook_cls, explicit, link, policy_cls):
+ error = _truncated_read(explicit=explicit)
+
+ prompt = _prompt_for(mock_hook_cls, policy_cls, error,
include_traceback=True)
+
+ assert prompt.startswith(f"{PROMPT_HEADER}Traceback (most recent call
last):\n")
+ assert "ConnectionResetError: peer closed connection after 25 of 4096
bytes" in prompt
+ assert link in prompt
+ assert prompt.endswith(f"json.decoder.JSONDecodeError: {error}")
+
+ @patch(HOOK, autospec=True)
+ def test_traceback_of_an_exception_never_raised_is_its_last_line(self,
mock_hook_cls):
+ prompt = _prompt_for(
+ mock_hook_cls, ClassifierRetryPolicy, ValueError("bad column
type"), include_traceback=True
+ )
+
+ assert prompt == f"{PROMPT_HEADER}ValueError: bad column type"
+
+ @patch(HOOK, autospec=True)
+ def test_redactor_gets_the_whole_traceback(self, mock_hook_cls):
+ secret = "s3cr3t-token"
+ redactor = create_autospec(
+ redact_registered_secrets, side_effect=lambda text:
text.replace(secret, "***")
+ )
+ error = _chained(RuntimeError("upload failed"),
PermissionError(f"token {secret} rejected"))
+
+ prompt = _prompt_for(
+ mock_hook_cls, ClassifierRetryPolicy, error,
include_traceback=True, redactor=redactor
+ )
+
+
redactor.assert_called_once_with("".join(traceback.format_exception(error)).rstrip("\n"))
+ assert secret not in prompt
+ assert "PermissionError: token *** rejected" in prompt
+ assert prompt.endswith("RuntimeError: upload failed")
+
+ @pytest.mark.parametrize(
+ "truncated", [pytest.param(True, id="over"), pytest.param(False,
id="exact-fit")]
+ )
+ @patch(HOOK, autospec=True)
+ def test_long_traceback_keeps_its_tail(self, mock_hook_cls, truncated):
+ error = _truncated_read()
+ full_text = "".join(traceback.format_exception(error)).rstrip("\n")
+ final_line = f"json.decoder.JSONDecodeError: {error}"
+ limit = len(final_line) if truncated else len(full_text)
+
+ prompt = _prompt_for(
+ mock_hook_cls, ClassifierRetryPolicy, error,
include_traceback=True, max_exception_length=limit
+ )
+
+ expected = f"(truncated) ...{final_line}" if truncated else full_text
+ assert prompt == f"{PROMPT_HEADER}{expected}"
+
class TestFallbackBehaviour:
"""When the LLM call itself fails the deterministic path decides,
unchanged."""