diff --git a/providers/common/ai/docs/retry_policies.rst b/providers/common/ai/docs/retry_policies.rst index 79625373783e5..686eb5c7e0f78 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 6440dcbaec658..cccdfbe45933b 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 @@ 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 @@ 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 @@ -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}") @@ -299,6 +317,7 @@ 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 @@ -306,14 +325,23 @@ def _hook(self) -> 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( @@ -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, @@ -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( 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 30ca9d2a679c1..bd34472b8de72 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 @@ def test_confidence_without_a_bar_is_recorded_but_not_gated(self, mock_hook_cls) 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 @@ def test_truncation_happens_after_redaction(self, mock_hook_cls): 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."""