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
99 changes: 87 additions & 12 deletions providers/common/ai/docs/retry_policies.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
----------------
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
from __future__ import annotations

import logging
import traceback
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import timedelta
Expand Down Expand Up @@ -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 --
Expand All @@ -238,17 +241,31 @@ def redact_registered_secrets(message: str) -> str:
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
Expand Down Expand Up @@ -279,6 +296,7 @@ def __init__(
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}")
Expand All @@ -299,21 +317,31 @@ def __init__(
)
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

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(
Expand Down Expand Up @@ -498,6 +526,7 @@ def __init__(
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,
Expand All @@ -508,6 +537,7 @@ def __init__(
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(
Expand Down
Loading