From efec91e59b05dea8beb0be677d4eb2f99bb90a23 Mon Sep 17 00:00:00 2001 From: Kaxil Naik Date: Wed, 30 Sep 2026 20:50:52 +0100 Subject: [PATCH 1/3] Support Airflow 2.11 in the Common AI provider Lower the provider's floor from Airflow 3.0 to 2.11, the floor the compat, standard and common-sql providers carry. The operators, decorators, hooks and toolsets run on 2.11 as they do on 3.0; features that need Airflow 3.1 or 3.3 (approval gates, HITL review, retry policies, tool approval, the task state store) keep their existing gates. - common.compat: export SET_DURING_EXECUTION, backed on Airflow 2 by an ArgNotSet that renders as Airflow 3's sentinel does, and resolve get_current_context through the standard provider first on Airflow 2, so it raises RuntimeError outside a task as Airflow 3 does. - PydanticAIHook.get_hook accepts hook_params on Airflow 2, mirroring Airflow 3. - The agent's per-attempt run key falls back to dag/run/task/map/try where the task instance has no id; spans then omit airflow.task_instance.id. - Declare structlog, which the provider imports but Airflow 2 does not ship, and route its output through the airflow.task logger on Airflow 2. - Declare the HITL review extra link only on Airflow 3.1+. - Run the provider's tests in the Airflow 2.11 compatibility job. --- .../src/airflow_breeze/global_constants.py | 2 +- providers/common/ai/README.rst | 3 +- providers/common/ai/docs/index.rst | 5 +- providers/common/ai/docs/installation.rst | 33 +++++++- providers/common/ai/docs/observability.rst | 3 + .../common/ai/docs/operators/llm_batch.rst | 4 +- providers/common/ai/docs/quickstart.rst | 3 +- .../common/ai/docs/self_hosted_models.rst | 3 +- providers/common/ai/pyproject.toml | 7 +- .../airflow/providers/common/ai/__init__.py | 4 +- .../providers/common/ai/batch/anthropic.py | 5 +- .../providers/common/ai/batch/openai.py | 5 +- .../providers/common/ai/batch/results.py | 5 +- .../providers/common/ai/decorators/agent.py | 2 +- .../providers/common/ai/decorators/llm.py | 2 +- .../common/ai/decorators/llm_batch.py | 2 +- .../common/ai/decorators/llm_branch.py | 2 +- .../common/ai/decorators/llm_file_analysis.py | 2 +- .../ai/decorators/llm_schema_compare.py | 2 +- .../providers/common/ai/decorators/llm_sql.py | 2 +- .../common/ai/durable/caching_model.py | 4 +- .../common/ai/durable/caching_toolset.py | 4 +- .../common/ai/durable/fingerprint.py | 4 +- .../providers/common/ai/durable/storage.py | 7 +- .../common/ai/durable/task_state_store.py | 4 +- .../providers/common/ai/hooks/pydantic_ai.py | 10 +++ .../providers/common/ai/observability.py | 19 ++++- .../providers/common/ai/operators/agent.py | 16 ++-- .../ai/operators/llamaindex_embedding.py | 2 +- .../ai/operators/llamaindex_retrieval.py | 2 +- .../providers/common/ai/utils/task_logger.py | 67 +++++++++++++++ .../providers/common/ai/utils/usage_budget.py | 5 +- .../unit/common/ai/batch/test_dispatch.py | 2 +- .../unit/common/ai/batch/test_results.py | 2 +- .../tests/unit/common/ai/batch/test_state.py | 2 +- .../common/ai/decorators/test_llm_batch.py | 3 +- .../common/ai/durable/test_replay_cost.py | 2 +- .../ai/durable/test_replay_verification.py | 2 +- .../unit/common/ai/durable/test_storage.py | 2 +- .../unit/common/ai/hooks/test_pydantic_ai.py | 15 ++++ .../unit/common/ai/mixins/test_approval.py | 2 +- .../unit/common/ai/operators/test_agent.py | 32 +++++++- .../ai/operators/test_document_loader.py | 12 +-- .../ai/operators/test_llamaindex_embedding.py | 2 +- .../ai/operators/test_llamaindex_retrieval.py | 4 +- .../common/ai/operators/test_llm_batch.py | 3 +- .../unit/common/ai/test_observability.py | 19 +++++ .../common/ai/toolsets/test_object_storage.py | 8 +- .../unit/common/ai/utils/test_task_logger.py | 82 +++++++++++++++++++ .../common/compat/_set_during_execution.py | 40 +++++++++ .../airflow/providers/common/compat/sdk.py | 14 +++- .../tests/unit/common/compat/test_sdk.py | 27 ++++-- uv.lock | 2 + 53 files changed, 429 insertions(+), 88 deletions(-) create mode 100644 providers/common/ai/src/airflow/providers/common/ai/utils/task_logger.py create mode 100644 providers/common/ai/tests/unit/common/ai/utils/test_task_logger.py create mode 100644 providers/common/compat/src/airflow/providers/common/compat/_set_during_execution.py diff --git a/dev/breeze/src/airflow_breeze/global_constants.py b/dev/breeze/src/airflow_breeze/global_constants.py index 43dfcf7e7ed93..6befa6a056621 100644 --- a/dev/breeze/src/airflow_breeze/global_constants.py +++ b/dev/breeze/src/airflow_breeze/global_constants.py @@ -862,7 +862,7 @@ def get_airflow_extras(): { "python-version": "3.10", "airflow-version": "2.11.1", - "remove-providers": "anthropic common.messaging common.dataquality edge3 fab git keycloak informatica common.ai modal opensearch", + "remove-providers": "anthropic common.messaging common.dataquality edge3 fab git keycloak informatica modal opensearch", "run-unit-tests": "true", }, { diff --git a/providers/common/ai/README.rst b/providers/common/ai/README.rst index 9d196fafb8d9a..b1204abc6c85f 100644 --- a/providers/common/ai/README.rst +++ b/providers/common/ai/README.rst @@ -53,10 +53,11 @@ Requirements ========================================== ================== PIP package Version required ========================================== ================== -``apache-airflow`` ``>=3.0.0`` +``apache-airflow`` ``>=2.11.0`` ``apache-airflow-providers-common-compat`` ``>=1.15.0`` ``apache-airflow-providers-standard`` ``>=1.20.0`` ``pydantic-ai-slim`` ``>=2.33.0`` +``structlog`` ``>=24.2.0`` ========================================== ================== Optional cross provider package dependencies diff --git a/providers/common/ai/docs/index.rst b/providers/common/ai/docs/index.rst index 666450bafccf5..fb2d2d62caa14 100644 --- a/providers/common/ai/docs/index.rst +++ b/providers/common/ai/docs/index.rst @@ -172,15 +172,16 @@ For the minimum Airflow version supported, see ``Requirements`` below. Requirements ------------ -The minimum Apache Airflow version supported by this provider distribution is ``3.0.0``. +The minimum Apache Airflow version supported by this provider distribution is ``2.11.0``. ========================================== ================== PIP package Version required ========================================== ================== -``apache-airflow`` ``>=3.0.0`` +``apache-airflow`` ``>=2.11.0`` ``apache-airflow-providers-common-compat`` ``>=1.15.0`` ``apache-airflow-providers-standard`` ``>=1.20.0`` ``pydantic-ai-slim`` ``>=2.33.0`` +``structlog`` ``>=24.2.0`` ========================================== ================== Optional cross provider package dependencies diff --git a/providers/common/ai/docs/installation.rst b/providers/common/ai/docs/installation.rst index 43e12ca82c83b..1f8c61086131e 100644 --- a/providers/common/ai/docs/installation.rst +++ b/providers/common/ai/docs/installation.rst @@ -20,7 +20,7 @@ Installation ============ -The provider needs Airflow 3.0 or later. Install it with the extra that matches the +The provider needs Airflow 2.11 or later. Install it with the extra that matches the model vendor your connection will point at: .. code-block:: bash @@ -62,7 +62,7 @@ package each extra installs. Features gated on the Airflow version ------------------------------------- -The provider runs on Airflow 3.0, but some features need a newer core: +The provider runs on Airflow 2.11, but some features need a newer core: .. list-table:: :header-rows: 1 @@ -70,13 +70,42 @@ The provider runs on Airflow 3.0, but some features need a newer core: * - Feature - Needs + * - The ``skills`` and ``git`` extras (``apache-airflow-providers-git`` needs Airflow 3) + - Airflow 3.0 * - :doc:`Approval gates ` and :doc:`HITL review ` - Airflow 3.1 + * - The **Model** field in the connection form; on older cores put the model in + **Extra**, for example ``{"model": "openai:gpt-5"}`` + - Airflow 3.2 * - :doc:`Retry policies ` - Airflow 3.3 * - :doc:`Durable execution ` without configuring ``[common.ai] durable_cache_path`` (the task state store) - Airflow 3.3 + * - :doc:`Tool approval ` that pauses the task; on older cores a tool + marked for approval fails the task + - Airflow 3.3 + * - A :doc:`structured output ` reaching downstream tasks as the + Pydantic model; on older cores it arrives as a ``dict`` + - Airflow 3.3 + +Airflow 2 +--------- + +On Airflow 2.11 the operators, decorators, hooks and toolsets run as they do on Airflow +3.0, apart from the table above. Three things differ from an Airflow 3 install: + +* The examples in these docs import ``dag``, ``task`` and ``Param`` from ``airflow.sdk``. + On Airflow 2 import ``dag`` and ``task`` from ``airflow.decorators`` and ``Param`` from + ``airflow.models.param``; the provider's own imports stay the same. +* Install Airflow with its constraints file as usual, then add the provider without it. The + Airflow 2.11 constraints pin ``apache-airflow-providers-common-compat`` and + ``apache-airflow-providers-common-sql`` to releases older than this provider needs. + Installing Airflow 2.11.0 without its constraints can also pull in a ``universal-pathlib`` + 0.3 release, which Airflow 2's ``ObjectStoragePath`` rejects; 2.11.1 and later cap it. + Leave out the ``skills`` and ``git`` extras: they need Airflow 3, and without + constraints ``pip`` upgrades Airflow to satisfy them. +* Python 3.10 to 3.12: the provider needs 3.10 or later, and Airflow 2.11 supports up to 3.12. Next steps ---------- diff --git a/providers/common/ai/docs/observability.rst b/providers/common/ai/docs/observability.rst index 828b28d1dda6e..3a79c2c64b109 100644 --- a/providers/common/ai/docs/observability.rst +++ b/providers/common/ai/docs/observability.rst @@ -81,6 +81,9 @@ How it works tool approval (see :doc:`tool_approval`) continues as ``-resumed``, which is the ``run_id`` the operator pushes; ``usage`` covers both parts. + Airflow 2 has no task-instance id, so there the key is + ``////``, and spans carry the five + identity keys without ``airflow.task_instance.id``. * **Scope.** The ``run_id`` / ``usage`` XComs come only from ``AgentOperator`` and ``@task.agent``, and so do the ``airflow.*`` identity attributes, apart from a Strands or ADK agent run inside ``agent_framework_tracing`` (see below). The other LLM diff --git a/providers/common/ai/docs/operators/llm_batch.rst b/providers/common/ai/docs/operators/llm_batch.rst index 806517c4051ad..c9030bf729711 100644 --- a/providers/common/ai/docs/operators/llm_batch.rst +++ b/providers/common/ai/docs/operators/llm_batch.rst @@ -282,8 +282,8 @@ provider's own batch listing first. ``cancel_on_kill`` cancels the batch if the task is killed. In deferrable mode this runs from the trigger's ``on_kill``, which only **Airflow 3.3+** calls; on those versions clearing, marking success or marking failed on a deferred task from the UI counts as a kill, so the batch is -cancelled and the next attempt submits a fresh one rather than re-attaching. On Airflow 3.0 to -3.2 a killed deferred task's batch keeps running and a clear re-attaches to it. Set +cancelled and the next attempt submits a fresh one rather than re-attaching. Before +Airflow 3.3 a killed deferred task's batch keeps running and a clear re-attaches to it. Set ``cancel_on_kill=False`` if you want clear-to-re-attach on 3.3+ as well. ``cancel_on_timeout=False`` lets a batch keep running (and billing) past this task's own diff --git a/providers/common/ai/docs/quickstart.rst b/providers/common/ai/docs/quickstart.rst index a21393ed190ba..c32379cc4e074 100644 --- a/providers/common/ai/docs/quickstart.rst +++ b/providers/common/ai/docs/quickstart.rst @@ -25,7 +25,8 @@ which one task asks a model to summarize release notes and a second task uses th At the end you know where the model's output lands and what a successful run looks like. You need a working :doc:`Airflow installation ` on -Airflow 3.0 or later and an API key for the model vendor you plan to use. Step 4 makes one +Airflow 2.11 or later and an API key for the model vendor you plan to use. On Airflow 2, +see :ref:`howto/installation` for what differs. Step 4 makes one real, billed API call. 1. Install the provider diff --git a/providers/common/ai/docs/self_hosted_models.rst b/providers/common/ai/docs/self_hosted_models.rst index 5eda43b7a3bac..ef1244f78f290 100644 --- a/providers/common/ai/docs/self_hosted_models.rst +++ b/providers/common/ai/docs/self_hosted_models.rst @@ -32,7 +32,8 @@ Before you start ------------------ This guide assumes a working :doc:`apache-airflow:installation/index` -(Airflow 3.0+) already exists. Its job stops at wiring Airflow to a server +(Airflow 2.11+; on Airflow 2 see :ref:`howto/installation` for +what differs) already exists. Its job stops at wiring Airflow to a server that's already running -- it doesn't cover installing or operating the model-serving stack itself. diff --git a/providers/common/ai/pyproject.toml b/providers/common/ai/pyproject.toml index dd818c99ad981..18d50b3033e78 100644 --- a/providers/common/ai/pyproject.toml +++ b/providers/common/ai/pyproject.toml @@ -67,14 +67,17 @@ requires-python = ">=3.10" # Make sure to run ``prek update-providers-dependencies --all-files`` # After you modify the dependencies, and rebuild your Breeze CI image with ``breeze ci-image build`` dependencies = [ - "apache-airflow>=3.0.0", - "apache-airflow-providers-common-compat>=1.15.0", + "apache-airflow>=2.11.0", + "apache-airflow-providers-common-compat>=1.15.0", # use next version "apache-airflow-providers-standard>=1.20.0", # 2.33.0 is the first release that works with anthropic>=1: it moved to httpx2 alongside # the SDK and stopped passing temperature/top_p/top_k as messages.create() kwargs, both # of which raise TypeError on 2.31 and earlier. The cost API this provider relies on # (RunUsage.cost, UsageLimits.cost_limit) landed earlier, in 2.23.0. "pydantic-ai-slim>=2.33.0", + # Airflow 3 brings structlog in through the Task SDK; Airflow 2 does not. 24.2.0 is the first + # release whose render_to_log_kwargs hands ``stacklevel`` to stdlib logging under that name. + "structlog>=24.2.0", ] # The optional dependencies should be modified in place in the generated file diff --git a/providers/common/ai/src/airflow/providers/common/ai/__init__.py b/providers/common/ai/src/airflow/providers/common/ai/__init__.py index db65dd598255d..a9ebbadf0fd02 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/__init__.py +++ b/providers/common/ai/src/airflow/providers/common/ai/__init__.py @@ -32,8 +32,8 @@ __version__ = "0.10.0" if packaging.version.parse(packaging.version.parse(airflow_version).base_version) < packaging.version.parse( - "3.0.0" + "2.11.0" ): raise RuntimeError( - f"The package `apache-airflow-providers-common-ai:{__version__}` needs Apache Airflow 3.0.0+" + f"The package `apache-airflow-providers-common-ai:{__version__}` needs Apache Airflow 2.11.0+" ) diff --git a/providers/common/ai/src/airflow/providers/common/ai/batch/anthropic.py b/providers/common/ai/src/airflow/providers/common/ai/batch/anthropic.py index 1fc58b45aeb50..d928026dddba9 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/batch/anthropic.py +++ b/providers/common/ai/src/airflow/providers/common/ai/batch/anthropic.py @@ -32,8 +32,6 @@ from collections.abc import Iterator from typing import TYPE_CHECKING, Any -import structlog - from airflow.providers.common.ai.batch.base import ( BatchAdapter, BatchState, @@ -43,9 +41,10 @@ SubmitResult, ) from airflow.providers.common.ai.exceptions import LLMBatchLimitExceededError +from airflow.providers.common.ai.utils.task_logger import get_task_logger from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException -log = structlog.get_logger(logger_name="task") +log = get_task_logger() if TYPE_CHECKING: from airflow.providers.common.ai.batch.base import BatchRequest diff --git a/providers/common/ai/src/airflow/providers/common/ai/batch/openai.py b/providers/common/ai/src/airflow/providers/common/ai/batch/openai.py index 11273789c3a6d..323ff9062f52b 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/batch/openai.py +++ b/providers/common/ai/src/airflow/providers/common/ai/batch/openai.py @@ -32,8 +32,6 @@ from datetime import datetime from typing import TYPE_CHECKING, Any -import structlog - from airflow.providers.common.ai.batch.base import ( BatchAdapter, BatchState, @@ -43,9 +41,10 @@ SubmitResult, ) from airflow.providers.common.ai.exceptions import LLMBatchLimitExceededError, LLMBatchModelMismatchError +from airflow.providers.common.ai.utils.task_logger import get_task_logger from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException -log = structlog.get_logger(logger_name="task") +log = get_task_logger() if TYPE_CHECKING: from airflow.providers.common.ai.batch.base import BatchRequest diff --git a/providers/common/ai/src/airflow/providers/common/ai/batch/results.py b/providers/common/ai/src/airflow/providers/common/ai/batch/results.py index 97e771d656763..3afba5d23bbf2 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/batch/results.py +++ b/providers/common/ai/src/airflow/providers/common/ai/batch/results.py @@ -34,17 +34,16 @@ from dataclasses import dataclass from typing import TYPE_CHECKING, Any -import structlog - from airflow.providers.common.ai.batch.base import evaluate_batch_counts from airflow.providers.common.ai.batch.output_schema import validate_extracted_output +from airflow.providers.common.ai.utils.task_logger import get_task_logger if TYPE_CHECKING: from airflow.providers.common.ai.batch.base import BatchAdapter, RawResultItem from airflow.providers.common.ai.batch.output_schema import OutputSpec from airflow.sdk import ObjectStoragePath -log = structlog.get_logger(logger_name="task") +log = get_task_logger() STATUS_SUCCESS = "success" STATUS_ERROR = "error" diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py index 25272d060092f..d2a92e076685a 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py @@ -33,13 +33,13 @@ validate_prompt, ) from airflow.providers.common.compat.sdk import ( + SET_DURING_EXECUTION, DecoratedOperator, TaskDecorator, context_merge, determine_kwargs, task_decorator_factory, ) -from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION if TYPE_CHECKING: from airflow.sdk import Context diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm.py index 8f29226202ab2..5e3ec84c6d1b5 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm.py @@ -35,13 +35,13 @@ validate_prompt, ) from airflow.providers.common.compat.sdk import ( + SET_DURING_EXECUTION, DecoratedOperator, TaskDecorator, context_merge, determine_kwargs, task_decorator_factory, ) -from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION if TYPE_CHECKING: from airflow.sdk import Context diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_batch.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_batch.py index d7f351b3728ef..d1bc139a1b498 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_batch.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_batch.py @@ -29,13 +29,13 @@ from airflow.providers.common.ai.operators.llm_batch import LLMBatchOperator from airflow.providers.common.compat.sdk import ( + SET_DURING_EXECUTION, DecoratedOperator, TaskDecorator, context_merge, determine_kwargs, task_decorator_factory, ) -from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION if TYPE_CHECKING: from airflow.sdk import Context diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_branch.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_branch.py index d835feeab17a1..fd8ee3f27279e 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_branch.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_branch.py @@ -33,13 +33,13 @@ validate_prompt, ) from airflow.providers.common.compat.sdk import ( + SET_DURING_EXECUTION, DecoratedOperator, TaskDecorator, context_merge, determine_kwargs, task_decorator_factory, ) -from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION if TYPE_CHECKING: from airflow.sdk import Context diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_file_analysis.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_file_analysis.py index 305184fb0d102..83447e9232266 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_file_analysis.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_file_analysis.py @@ -23,13 +23,13 @@ from airflow.providers.common.ai.operators.llm_file_analysis import LLMFileAnalysisOperator from airflow.providers.common.compat.sdk import ( + SET_DURING_EXECUTION, DecoratedOperator, TaskDecorator, context_merge, determine_kwargs, task_decorator_factory, ) -from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION if TYPE_CHECKING: from airflow.sdk import Context diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_schema_compare.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_schema_compare.py index 2e106509215e8..a4f22bfa48674 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_schema_compare.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_schema_compare.py @@ -33,13 +33,13 @@ validate_prompt, ) from airflow.providers.common.compat.sdk import ( + SET_DURING_EXECUTION, DecoratedOperator, TaskDecorator, context_merge, determine_kwargs, task_decorator_factory, ) -from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION if TYPE_CHECKING: from airflow.sdk import Context diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_sql.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_sql.py index 7ea9765ddcdbb..e06a46183b8c7 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_sql.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/llm_sql.py @@ -33,13 +33,13 @@ validate_prompt, ) from airflow.providers.common.compat.sdk import ( + SET_DURING_EXECUTION, DecoratedOperator, TaskDecorator, context_merge, determine_kwargs, task_decorator_factory, ) -from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION if TYPE_CHECKING: from airflow.sdk import Context diff --git a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py index ef1d6940f9c3e..07084d5fdbd4c 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py +++ b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py @@ -21,14 +21,14 @@ from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any -import structlog from pydantic_ai.messages import ModelResponse, ToolCallPart from pydantic_ai.models.wrapper import WrapperModel from airflow.providers.common.ai.durable.base import build_model_step_key, build_tool_step_key from airflow.providers.common.ai.durable.fingerprint import fingerprint_model_request +from airflow.providers.common.ai.utils.task_logger import get_task_logger -log = structlog.get_logger(logger_name="task") +log = get_task_logger() if TYPE_CHECKING: from pydantic_ai.messages import ModelMessage diff --git a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py index fa4a61b9af74b..1eb30127c01ff 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py +++ b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py @@ -21,11 +21,11 @@ from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any -import structlog from pydantic_ai.toolsets.wrapper import WrapperToolset from airflow.providers.common.ai.durable.base import build_tool_step_key from airflow.providers.common.ai.durable.fingerprint import fingerprint_tool_call +from airflow.providers.common.ai.utils.task_logger import get_task_logger from airflow.providers.common.ai.utils.tool_metrics import record_tool_call from airflow.providers.common.ai.utils.toolset_base import AirflowToolset @@ -36,7 +36,7 @@ from airflow.providers.common.ai.durable.replay_usage import ReplayUsageLedger from airflow.providers.common.ai.durable.step_counter import DurableStepCounter -log = structlog.get_logger(logger_name="task") +log = get_task_logger() @dataclass diff --git a/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py b/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py index bc4a291ea929a..be3c2b6f1c212 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py +++ b/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py @@ -45,18 +45,18 @@ import json from typing import TYPE_CHECKING, Any -import structlog from pydantic import TypeAdapter from pydantic_ai.messages import ModelMessagesTypeAdapter from pydantic_ai.models import ModelRequestParameters from airflow.providers.common.ai.utils.prompt_cache import PROMPT_CACHE_SETTING_NAMES +from airflow.providers.common.ai.utils.task_logger import get_task_logger if TYPE_CHECKING: from pydantic_ai.messages import ModelMessage from pydantic_ai.settings import ModelSettings -log = structlog.get_logger(logger_name="task") +log = get_task_logger() _MODEL_REQUEST_PARAMETERS_ADAPTER = TypeAdapter(ModelRequestParameters) diff --git a/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py b/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py index a3f974bcb4240..4e98fe40eb47c 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py +++ b/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py @@ -24,22 +24,21 @@ from functools import lru_cache from typing import Any -import structlog from pydantic_ai.messages import ModelMessagesTypeAdapter, ModelResponse # Sentinel to distinguish "cached None" from "no cache entry" for tool results. # Shared with the task state store backend so the envelope shape cannot drift. from airflow.providers.common.ai.durable.base import TOOL_RESULT_SENTINEL as _SENTINEL +from airflow.providers.common.ai.utils.task_logger import get_task_logger -log = structlog.get_logger(logger_name="task") +log = get_task_logger() SECTION = "common.ai" @lru_cache(maxsize=1) def _get_base_path(): - from airflow.providers.common.compat.sdk import conf - from airflow.sdk import ObjectStoragePath + from airflow.providers.common.compat.sdk import ObjectStoragePath, conf path = conf.get(SECTION, "durable_cache_path", fallback="") if not path: diff --git a/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py b/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py index ecf1c0e066d8c..c02a33881c1c5 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py +++ b/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py @@ -37,10 +37,10 @@ import json from typing import TYPE_CHECKING, Any -import structlog from pydantic_ai.messages import ModelMessagesTypeAdapter from airflow.providers.common.ai.durable.base import TOOL_RESULT_SENTINEL +from airflow.providers.common.ai.utils.task_logger import get_task_logger from airflow.sdk.execution_time.context import NEVER_EXPIRE if TYPE_CHECKING: @@ -48,7 +48,7 @@ from airflow.sdk.execution_time.context import TaskStateStoreAccessor -log = structlog.get_logger(logger_name="task") +log = get_task_logger() class TaskStateStoreDurableStorage: diff --git a/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py b/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py index c8531d4fe125b..efb34edaac0e1 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py +++ b/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py @@ -158,6 +158,16 @@ def __init__( self._conn: Connection | None = None self._conn_extra_dejson: dict[str, Any] = {} + @classmethod + def get_hook(cls, conn_id: str, hook_params: dict | None = None): + """ + Return the hook for ``conn_id``, built with ``hook_params``. + + Airflow 3's ``BaseHook.get_hook`` already takes ``hook_params``; Airflow 2's does + not, so this mirrors the Airflow 3 body. + """ + return cls.get_connection(conn_id).get_hook(hook_params=hook_params) + @staticmethod def get_ui_field_behaviour() -> dict[str, Any]: """Return custom field behaviour for the Airflow connection form.""" diff --git a/providers/common/ai/src/airflow/providers/common/ai/observability.py b/providers/common/ai/src/airflow/providers/common/ai/observability.py index 25e4a8712b781..0d70687b28e43 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/observability.py +++ b/providers/common/ai/src/airflow/providers/common/ai/observability.py @@ -121,15 +121,30 @@ def build_run_identity_attributes(ti: Any) -> dict[str, Any]: Reuses core's task-span attribute keys (see ``_make_task_span``) so agent spans filter identically to the task span they nest under, plus the per-attempt task-instance id as the run join key carried on every span. + Airflow 2 task instances have no id, so the attribute is left out there. """ - return { + attributes: dict[str, Any] = { "airflow.dag_id": ti.dag_id, "airflow.task_id": ti.task_id, "airflow.dag_run.run_id": ti.run_id, "airflow.task_instance.try_number": ti.try_number, "airflow.task_instance.map_index": ti.map_index if ti.map_index is not None else -1, - "airflow.task_instance.id": str(ti.id), } + if (ti_id := getattr(ti, "id", None)) is not None: + attributes["airflow.task_instance.id"] = str(ti_id) + return attributes + + +def task_instance_run_key(ti: Any) -> str: + """ + Return a per-attempt key for ``ti``: its id on Airflow 3, a composite on Airflow 2. + + Airflow 2 task instances have no ``id`` column; dag, run, task, map index and try + number identify one attempt just as uniquely. + """ + if (ti_id := getattr(ti, "id", None)) is not None: + return str(ti_id) + return f"{ti.dag_id}/{ti.run_id}/{ti.task_id}/{ti.map_index}/{ti.try_number}" def stamp_identity_on_agent_spans(agent: Agent, attributes: dict[str, Any]) -> None: diff --git a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py index 974c7697cdd05..09e464844e939 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py @@ -48,6 +48,7 @@ from airflow.providers.common.ai.observability import ( build_run_identity_attributes, stamp_identity_on_agent_spans, + task_instance_run_key, ) from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset from airflow.providers.common.ai.utils.logging import ( @@ -472,7 +473,9 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): "usage_limits", ) - operator_extra_links = (HITLReviewLink(),) + # HITL review needs Airflow 3.1. Airflow 2 would also log an error for the unregistered + # link class every time the webserver loads a Dag with this operator. + operator_extra_links = (HITLReviewLink(),) if AIRFLOW_V_3_1_PLUS else () def __init__( self, @@ -996,7 +999,7 @@ def _report_failed_run(self, context: Context, run_usage: RunUsage) -> None: return ti = context["task_instance"] try: - ti.xcom_push(key="run_id", value=str(ti.id)) + ti.xcom_push(key="run_id", value=task_instance_run_key(ti)) except Exception: self.log.warning("Failed to push run_id XCom for the failed run", exc_info=True) if attempt_usage is not None: @@ -1106,10 +1109,11 @@ def execute(self, context: Context) -> Any: self._run_identity_attrs = build_run_identity_attributes(ti) stamp_identity_on_agent_spans(agent, self._run_identity_attrs) - # The task-instance id is non-nullable and regenerated on each retry, so it - # is a unique, reverse-resolvable join key. It lands on result.run_id, the - # run's messages, and the ``gen_ai.agent.call.id`` span attribute. - run_kwargs: dict[str, Any] = {"usage_limits": usage_limits, "run_id": str(ti.id)} + # A per-attempt key (the task-instance id on Airflow 3, which is regenerated on + # each retry; dag/run/task/map/try on Airflow 2) is a unique, reverse-resolvable + # join key. It lands on result.run_id, the run's messages, and the + # ``gen_ai.agent.call.id`` span attribute. + run_kwargs: dict[str, Any] = {"usage_limits": usage_limits, "run_id": task_instance_run_key(ti)} history = self._resolve_message_history() if history is not None: run_kwargs["message_history"] = history diff --git a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py index c7003db1e3343..8a878a874f766 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py +++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py @@ -228,7 +228,7 @@ def _resolve_embed_model(self) -> BaseEmbedding: def _persist(self, index: Any, persist_dir: str) -> None: """Persist the index to ``persist_dir``; cloud URIs go through ObjectStoragePath.""" if "://" in persist_dir: - from airflow.sdk import ObjectStoragePath + from airflow.providers.common.compat.sdk import ObjectStoragePath target = ObjectStoragePath(persist_dir, conn_id=self.persist_conn_id) target.mkdir(parents=True, exist_ok=True) diff --git a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py index 234a3bfdd4d1f..d61e1c5f2bef0 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py +++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py @@ -193,7 +193,7 @@ def _resolve_embed_model(self) -> BaseEmbedding: def _open_storage_context(self, storage_context_cls: Any) -> Any: """Open a ``StorageContext`` from a local path or storage URI.""" if "://" in self.index_persist_dir: - from airflow.sdk import ObjectStoragePath + from airflow.providers.common.compat.sdk import ObjectStoragePath source = ObjectStoragePath(self.index_persist_dir, conn_id=self.persist_conn_id) if not source.is_dir(): diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/task_logger.py b/providers/common/ai/src/airflow/providers/common/ai/utils/task_logger.py new file mode 100644 index 0000000000000..5bb9c07bef284 --- /dev/null +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/task_logger.py @@ -0,0 +1,67 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""A structlog logger that writes to the task log on Airflow 2 and Airflow 3.""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Any + +import structlog + +from airflow.providers.common.compat.version_compat import AIRFLOW_V_3_0_PLUS + +if TYPE_CHECKING: + from structlog.typing import EventDict, FilteringBoundLogger + +_STDLIB_LOG_KWARGS = ("exc_info", "stack_info", "stacklevel") +# Frames between a ``log.warning(...)`` call and the stdlib logger inside structlog's +# BoundLogger, so the record names the caller's file and line rather than structlog's. +_CALLER_STACKLEVEL = 4 + + +def _fold_fields_into_message(_logger: Any, _method_name: str, event_dict: EventDict) -> EventDict: + """Append ``key=value`` fields to the message, since Airflow 2's task log format drops ``extra``.""" + fields = [key for key in event_dict if key != "event" and key not in _STDLIB_LOG_KWARGS] + if fields: + rendered = " ".join(f"{key}={event_dict.pop(key)!r}" for key in fields) + event_dict["event"] = f"{event_dict['event']} {rendered}" + event_dict.setdefault("stacklevel", _CALLER_STACKLEVEL) + return event_dict + + +def get_task_logger() -> FilteringBoundLogger: + """ + Return a structlog logger that writes to the task log. + + Airflow 3 configures structlog for task processes. Airflow 2 does not, so structlog there + falls back to its defaults and prints every level, ``debug`` included, to stdout. On Airflow 2 the + logger wraps the ``airflow.task`` stdlib logger instead, which applies the task log's level + and handlers, without changing the process-wide structlog configuration. + """ + if AIRFLOW_V_3_0_PLUS: + return structlog.get_logger(logger_name="task") + return structlog.wrap_logger( + logging.getLogger("airflow.task"), + wrapper_class=structlog.stdlib.BoundLogger, + processors=[ + structlog.stdlib.filter_by_level, + structlog.stdlib.PositionalArgumentsFormatter(), + _fold_fields_into_message, + structlog.stdlib.render_to_log_kwargs, + ], + ) diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py b/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py index 9e2f1fb606711..f67c686fa99f2 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py @@ -32,13 +32,14 @@ from decimal import Decimal, InvalidOperation from typing import TYPE_CHECKING, Any -import structlog from pydantic_ai.usage import RunUsage +from airflow.providers.common.ai.utils.task_logger import get_task_logger + if TYPE_CHECKING: from airflow.sdk.execution_time.context import TaskStateStoreAccessor -log = structlog.get_logger(logger_name="task") +log = get_task_logger() # Reserved task state store key for the cumulative cross-attempt usage record. Separate # from durable's ``DURABLE_KEY_PREFIX`` (see durable/base.py) so it is never mistaken for diff --git a/providers/common/ai/tests/unit/common/ai/batch/test_dispatch.py b/providers/common/ai/tests/unit/common/ai/batch/test_dispatch.py index 430fa49ad467c..1603f6b2c05f9 100644 --- a/providers/common/ai/tests/unit/common/ai/batch/test_dispatch.py +++ b/providers/common/ai/tests/unit/common/ai/batch/test_dispatch.py @@ -31,7 +31,7 @@ BatchProviderNotYetSupportedError, UnsupportedBatchProviderError, ) -from airflow.sdk import Connection +from airflow.providers.common.compat.sdk import Connection class TestSplitModelId: diff --git a/providers/common/ai/tests/unit/common/ai/batch/test_results.py b/providers/common/ai/tests/unit/common/ai/batch/test_results.py index 82da64ac3a5e4..1cbb218a5a467 100644 --- a/providers/common/ai/tests/unit/common/ai/batch/test_results.py +++ b/providers/common/ai/tests/unit/common/ai/batch/test_results.py @@ -32,7 +32,7 @@ missing_indexes, stream_results_to_jsonl, ) -from airflow.sdk import ObjectStoragePath +from airflow.providers.common.compat.sdk import ObjectStoragePath class _FakeAdapter(BatchAdapter): diff --git a/providers/common/ai/tests/unit/common/ai/batch/test_state.py b/providers/common/ai/tests/unit/common/ai/batch/test_state.py index fb91d5f5fabf2..7b157f1e907bd 100644 --- a/providers/common/ai/tests/unit/common/ai/batch/test_state.py +++ b/providers/common/ai/tests/unit/common/ai/batch/test_state.py @@ -33,7 +33,7 @@ write_submitted, ) from airflow.providers.common.ai.exceptions import LLMBatchStateReadError -from airflow.sdk import ObjectStoragePath +from airflow.providers.common.compat.sdk import ObjectStoragePath @pytest.fixture diff --git a/providers/common/ai/tests/unit/common/ai/decorators/test_llm_batch.py b/providers/common/ai/tests/unit/common/ai/decorators/test_llm_batch.py index 538148c8faf80..f66bb87285e53 100644 --- a/providers/common/ai/tests/unit/common/ai/decorators/test_llm_batch.py +++ b/providers/common/ai/tests/unit/common/ai/decorators/test_llm_batch.py @@ -33,8 +33,7 @@ from airflow.providers.common.ai.decorators.llm_batch import _LLMBatchDecoratedOperator from airflow.providers.common.ai.exceptions import LLMBatchInputError from airflow.providers.common.ai.operators import llm_batch as llm_batch_module -from airflow.providers.common.compat.sdk import TaskDeferred -from airflow.sdk import DAG, Connection +from airflow.providers.common.compat.sdk import DAG, Connection, TaskDeferred class _FakeAdapter(BatchAdapter): diff --git a/providers/common/ai/tests/unit/common/ai/durable/test_replay_cost.py b/providers/common/ai/tests/unit/common/ai/durable/test_replay_cost.py index 3b083a2a0a177..4905d44c22aff 100644 --- a/providers/common/ai/tests/unit/common/ai/durable/test_replay_cost.py +++ b/providers/common/ai/tests/unit/common/ai/durable/test_replay_cost.py @@ -56,7 +56,7 @@ from airflow.providers.common.ai.durable.replay_usage import ReplayUsageLedger from airflow.providers.common.ai.durable.step_counter import DurableStepCounter from airflow.providers.common.ai.durable.storage import DurableStorage -from airflow.sdk import ObjectStoragePath +from airflow.providers.common.compat.sdk import ObjectStoragePath PRICED_COST = Decimal("0.10") diff --git a/providers/common/ai/tests/unit/common/ai/durable/test_replay_verification.py b/providers/common/ai/tests/unit/common/ai/durable/test_replay_verification.py index ad9c244b0c60f..1e57ad45eaa88 100644 --- a/providers/common/ai/tests/unit/common/ai/durable/test_replay_verification.py +++ b/providers/common/ai/tests/unit/common/ai/durable/test_replay_verification.py @@ -37,7 +37,7 @@ from airflow.providers.common.ai.durable.caching_toolset import CachingToolset from airflow.providers.common.ai.durable.step_counter import DurableStepCounter from airflow.providers.common.ai.durable.storage import DurableStorage -from airflow.sdk import ObjectStoragePath +from airflow.providers.common.compat.sdk import ObjectStoragePath @pytest.fixture diff --git a/providers/common/ai/tests/unit/common/ai/durable/test_storage.py b/providers/common/ai/tests/unit/common/ai/durable/test_storage.py index c1a4176d164c8..c2670655195f0 100644 --- a/providers/common/ai/tests/unit/common/ai/durable/test_storage.py +++ b/providers/common/ai/tests/unit/common/ai/durable/test_storage.py @@ -28,7 +28,7 @@ from pydantic_ai.usage import RequestUsage from airflow.providers.common.ai.durable.storage import DurableStorage -from airflow.sdk import ObjectStoragePath +from airflow.providers.common.compat.sdk import ObjectStoragePath @pytest.fixture diff --git a/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py b/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py index 2a4dc2b1fec6e..5f856c1d5c0b4 100644 --- a/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py +++ b/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py @@ -337,6 +337,21 @@ def registry(): yield reg +class TestPydanticAIHookGetHook: + def test_builds_the_connection_hook_with_hook_params(self, registry): + """Airflow 2's ``BaseHook.get_hook`` takes no ``hook_params``; the hook's own override does.""" + registry.add("llm") + + hook = PydanticAIHook.get_hook( + "llm", hook_params={"model_id": "openai:gpt-5", "fallback_conn_ids": []} + ) + + assert isinstance(hook, PydanticAIHook) + assert hook.llm_conn_id == "llm" + assert hook.model_id == "openai:gpt-5" + assert hook.fallback_conn_ids == [] + + class _InferModelStub: """Resolve every model string to its own recognisable model, and record how it was built.""" diff --git a/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py b/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py index 960999f0495c5..85123748ff60c 100644 --- a/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py +++ b/providers/common/ai/tests/unit/common/ai/mixins/test_approval.py @@ -35,8 +35,8 @@ LLMApprovalMixin, ) from airflow.providers.common.compat.notifier import BaseNotifier +from airflow.providers.common.compat.sdk import DAG from airflow.providers.standard.exceptions import HITLRejectException, HITLTriggerEventError -from airflow.sdk import DAG if AIRFLOW_V_3_3_PLUS: from airflow.sdk.exceptions import TaskAwaitingInput diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py index be384d21a0586..e344947308efe 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py @@ -86,11 +86,16 @@ copy_run_usage, dump_run_usage, ) -from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException, BaseHook -from airflow.sdk import DAG, task +from airflow.providers.common.compat.sdk import ( + DAG, + AirflowException, + AirflowOptionalProviderFeatureException, + BaseHook, + task, +) from tests_common.test_utils.compat import OperatorSerialization -from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS, AIRFLOW_V_3_3_PLUS +from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_1_PLUS, AIRFLOW_V_3_3_PLUS from unit.common.ai.sandbox.fake_tags import TaggedBackend try: @@ -248,7 +253,8 @@ def cleanup(self): class TestAgentOperatorValidation: def test_requires_llm_conn_id(self): - with pytest.raises(TypeError): + # Airflow 2's BaseOperator reports a missing required argument as AirflowException. + with pytest.raises(TypeError if AIRFLOW_V_3_0_PLUS else AirflowException): AgentOperator(task_id="test", prompt="hello") @pytest.mark.skipif( @@ -527,6 +533,10 @@ def __init__(self, wrapped, *, audit_name): pytest.param("decorator", "tenant_{{ task.op_kwargs.customer }}", id="decorator"), ], ) + @pytest.mark.skipif( + not AIRFLOW_V_3_0_PLUS, + reason="Airflow 2's MappedOperator resolves expansions through the metadata database and a task session", + ) def test_each_map_index_gets_its_own_connection(self, form, template): """Through the real MappedOperator render path, for both authoring forms.""" shared = SQLToolset(db_conn_id=template) @@ -773,6 +783,20 @@ def test_execute_creates_agent_from_hook(self, mock_hook_cls, make_mock_run_resu "What is the answer?", usage_limits=None, run_id="ti-1", cancellation_token=ANY, usage=ANY ) + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + def test_execute_keys_the_run_by_attempt_without_a_task_instance_id( + self, mock_hook_cls, make_mock_run_result + ): + """Airflow 2 task instances have no ``id``; the run key falls back to dag/run/task/map/try.""" + mock_agent = _make_mock_agent("ok", make_mock_run_result) + mock_hook_cls.get_hook.return_value.create_agent.return_value = mock_agent + op = AgentOperator(task_id="test", prompt="hello", llm_conn_id="my_llm") + + op.execute(context=_make_context(_make_ti(id=None, try_number=2))) + + _, kwargs = mock_agent.run_sync.call_args + assert kwargs["run_id"] == "dag/run/task/-1/2" + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) def test_execute_passes_toolsets_in_agent_kwargs(self, mock_hook_cls, make_mock_run_result): """Toolsets reach the agent wrapped for masking, then for logging.""" diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_document_loader.py b/providers/common/ai/tests/unit/common/ai/operators/test_document_loader.py index 60031e3edfb1a..dcd0ab8dac9ab 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_document_loader.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_document_loader.py @@ -23,7 +23,7 @@ import pytest from airflow.providers.common.ai.operators.document_loader import DocumentLoaderOperator -from airflow.sdk import DAG +from airflow.providers.common.compat.sdk import DAG class TestDocumentLoaderInit: @@ -528,7 +528,7 @@ def test_file_extensions_case_insensitive(self, tmp_path): class TestCloudUriDispatch: """``source_path`` containing a URI scheme routes through ObjectStoragePath.""" - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") def test_single_object_uri_returns_one_document(self, mock_osp_cls): # `str(mock_obj)` returns whatever MagicMock renders; we only assert # the file_name field, not file_path, so leaving __str__ default is @@ -552,7 +552,7 @@ def test_single_object_uri_returns_one_document(self, mock_osp_cls): assert result[0]["text"] == "cloud content" assert result[0]["metadata"]["file_name"] == "report.txt" - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") def test_directory_uri_iterates_children(self, mock_osp_cls): # Root is a directory; iterdir yields two text files. def _mock_child(name: str, content: bytes): @@ -577,7 +577,7 @@ def _mock_child(name: str, content: bytes): assert {doc["text"] for doc in result} == {"alpha", "beta"} - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") def test_neither_file_nor_dir_uri_raises(self, mock_osp_cls): bad = MagicMock() bad.is_file.return_value = False @@ -588,7 +588,7 @@ def test_neither_file_nor_dir_uri_raises(self, mock_osp_cls): with pytest.raises(FileNotFoundError, match="neither a file nor a directory"): op.execute(context=MagicMock()) - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") def test_glob_uri_matches_across_directories(self, mock_osp_cls): def _mock_match(name: str, content: bytes): match = MagicMock() @@ -613,7 +613,7 @@ def _mock_match(name: str, content: bytes): root.glob.assert_called_once_with("**/*.txt") assert {doc["text"] for doc in result} == {"alpha", "beta"} - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") def test_glob_in_bucket_segment_raises(self, mock_osp_cls): op = DocumentLoaderOperator(task_id="test", source_path="s3://bucket-*/dir/a.txt") with pytest.raises(ValueError, match="scheme or bucket segment"): diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py index e410dfded058f..0d647f372e089 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py @@ -287,7 +287,7 @@ def test_local_persist_dir_calls_makedirs_and_storage_persist( nodes_arg = _li["VectorStoreIndex"].call_args.args[0] assert nodes_arg[0].embedding == [0.1] - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") @patch("airflow.providers.common.ai.hooks.llamaindex.LlamaIndexHook.get_embedding_model") def test_cloud_uri_persist_dir_uses_object_storage_path(self, mock_get_embed, mock_osp_cls, _li): # ``ObjectStoragePath.__str__`` returns ``://@/...`` diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py index c0c4b0c9e8931..6c1378c5c6f57 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py @@ -201,7 +201,7 @@ def test_local_missing_dir_raises_with_hint(self, mock_get_embed, _li, tmp_path) with pytest.raises(FileNotFoundError, match="LlamaIndexEmbeddingOperator"): op.execute(context=MagicMock()) - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") @patch("airflow.providers.common.ai.hooks.llamaindex.LlamaIndexHook.get_embedding_model") def test_cloud_missing_uri_raises_with_hint(self, mock_get_embed, mock_osp_cls, _li): missing = MagicMock() @@ -219,7 +219,7 @@ def test_cloud_missing_uri_raises_with_hint(self, mock_get_embed, mock_osp_cls, class TestRetrievalOperatorCloudURI: - @patch("airflow.sdk.ObjectStoragePath") + @patch("airflow.providers.common.compat.sdk.ObjectStoragePath") @patch("airflow.providers.common.ai.hooks.llamaindex.LlamaIndexHook.get_embedding_model") def test_cloud_uri_opens_storage_with_fs(self, mock_get_embed, mock_osp_cls, _li): # ``ObjectStoragePath.__str__`` returns ``://@/...`` diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llm_batch.py b/providers/common/ai/tests/unit/common/ai/operators/test_llm_batch.py index f8703b9725b86..fbe2a04181987 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_llm_batch.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_llm_batch.py @@ -49,8 +49,7 @@ ) from airflow.providers.common.ai.operators import llm_batch as llm_batch_module from airflow.providers.common.ai.operators.llm_batch import LLMBatchOperator -from airflow.providers.common.compat.sdk import TaskDeferred -from airflow.sdk import Connection, ObjectStoragePath +from airflow.providers.common.compat.sdk import Connection, ObjectStoragePath, TaskDeferred class Diagnosis(BaseModel): diff --git a/providers/common/ai/tests/unit/common/ai/test_observability.py b/providers/common/ai/tests/unit/common/ai/test_observability.py index 91f4e8c91dfc0..6ba883e9b6d49 100644 --- a/providers/common/ai/tests/unit/common/ai/test_observability.py +++ b/providers/common/ai/tests/unit/common/ai/test_observability.py @@ -162,6 +162,25 @@ def test_builds_expected_attributes(self, map_index, expected_map_index): "airflow.task_instance.id": "ti-1", } + def test_leaves_out_the_task_instance_id_on_airflow_2(self): + """A composite key is not a task-instance id, so Airflow 2 spans carry the parts alone.""" + ti = SimpleNamespace(dag_id="d", task_id="t", run_id="r", try_number=2, map_index=-1) + + assert "airflow.task_instance.id" not in observability.build_run_identity_attributes(ti) + + +class TestTaskInstanceRunKey: + def test_uses_the_task_instance_id_when_it_has_one(self): + ti = SimpleNamespace(id="0199-uuid", dag_id="d", task_id="t", run_id="r", try_number=2, map_index=3) + + assert observability.task_instance_run_key(ti) == "0199-uuid" + + def test_builds_a_per_attempt_key_without_an_id(self): + """Airflow 2 task instances have no ``id`` column.""" + ti = SimpleNamespace(dag_id="d", task_id="t", run_id="r", try_number=2, map_index=3) + + assert observability.task_instance_run_key(ti) == "d/r/t/3/2" + class TestStampIdentityOnAgentSpans: _ATTRS = { diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_object_storage.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_object_storage.py index 989552db68f9c..0dc15cdd5b4ae 100644 --- a/providers/common/ai/tests/unit/common/ai/toolsets/test_object_storage.py +++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_object_storage.py @@ -32,7 +32,13 @@ from pydantic_ai.usage import RunUsage from airflow.providers.common.ai.toolsets.object_storage import ObjectStorageToolset -from airflow.sdk.io.store import _STORE_CACHE, ObjectStore + +from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS + +if AIRFLOW_V_3_0_PLUS: + from airflow.sdk.io.store import _STORE_CACHE, ObjectStore +else: + from airflow.io.store import _STORE_CACHE, ObjectStore # type: ignore[no-redef] @pytest.fixture diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_task_logger.py b/providers/common/ai/tests/unit/common/ai/utils/test_task_logger.py new file mode 100644 index 0000000000000..f48b0b6a7cdfa --- /dev/null +++ b/providers/common/ai/tests/unit/common/ai/utils/test_task_logger.py @@ -0,0 +1,82 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import logging +from unittest.mock import patch + +import pytest + +from airflow.providers.common.ai.utils import task_logger +from airflow.providers.common.ai.utils.task_logger import get_task_logger + + +class _RecordingHandler(logging.Handler): + def __init__(self) -> None: + super().__init__() + self.records: list[logging.LogRecord] = [] + + def emit(self, record: logging.LogRecord) -> None: + self.records.append(record) + + +@pytest.fixture +def airflow2_task_log(): + """Route ``get_task_logger`` down its Airflow 2 path and capture the ``airflow.task`` records.""" + airflow_task = logging.getLogger("airflow.task") + handler = _RecordingHandler() + previous_level = airflow_task.level + airflow_task.addHandler(handler) + airflow_task.setLevel(logging.INFO) + try: + with patch.object(task_logger, "AIRFLOW_V_3_0_PLUS", False): + yield handler.records + finally: + airflow_task.removeHandler(handler) + airflow_task.setLevel(previous_level) + + +class TestGetTaskLoggerOnAirflow2: + def test_applies_the_task_log_level(self, airflow2_task_log): + get_task_logger().debug("Durable: cached model response", step=0) + + assert airflow2_task_log == [] + + def test_folds_fields_into_the_message(self, airflow2_task_log): + get_task_logger().warning("Durable: cache miss", step=2, tool="get_weather") + + (record,) = airflow2_task_log + assert record.levelno == logging.WARNING + assert record.getMessage() == "Durable: cache miss step=2 tool='get_weather'" + + def test_names_the_calling_line_not_structlog(self, airflow2_task_log): + get_task_logger().warning("from the caller") + + (record,) = airflow2_task_log + assert record.pathname == __file__ + assert record.funcName == "test_names_the_calling_line_not_structlog" + + def test_hands_exc_info_to_stdlib(self, airflow2_task_log): + try: + raise ValueError("boom") + except ValueError: + get_task_logger().warning("Failed to write the cache", exc_info=True) + + (record,) = airflow2_task_log + assert record.getMessage() == "Failed to write the cache" + assert record.exc_info is not None + assert record.exc_info[0] is ValueError diff --git a/providers/common/compat/src/airflow/providers/common/compat/_set_during_execution.py b/providers/common/compat/src/airflow/providers/common/compat/_set_during_execution.py new file mode 100644 index 0000000000000..729fc3db16489 --- /dev/null +++ b/providers/common/compat/src/airflow/providers/common/compat/_set_during_execution.py @@ -0,0 +1,40 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +""" +Airflow 2 stand-in for the Task SDK's ``SET_DURING_EXECUTION`` sentinel. + +A decorated operator passes the sentinel for an argument its callable fills in, such as the prompt +of ``@task.llm``. Airflow 2 stores a template field it cannot JSON-encode as ``str(value)``, which +for a bare ``NOTSET`` is an object address: it differs in every process, so the serialized Dag's +hash and the rendered template view change with it. This sentinel renders the way Airflow 3's +does. :mod:`airflow.providers.common.compat.sdk` tries the SDK first, so on Airflow 3 this module +is never imported. +""" + +from __future__ import annotations + +from airflow.utils.types import ArgNotSet # type: ignore[attr-defined] # Airflow 2 only + + +class SetDuringExecution(ArgNotSet): + """Sentinel for an argument that is set during execution, not at parse time.""" + + def __repr__(self) -> str: + return "DYNAMIC (set during execution)" + + +SET_DURING_EXECUTION = SetDuringExecution() diff --git a/providers/common/compat/src/airflow/providers/common/compat/sdk.py b/providers/common/compat/src/airflow/providers/common/compat/sdk.py index de7480c733650..f37a194b337e8 100644 --- a/providers/common/compat/src/airflow/providers/common/compat/sdk.py +++ b/providers/common/compat/src/airflow/providers/common/compat/sdk.py @@ -83,6 +83,7 @@ from airflow.sdk.bases.sensor import poke_mode_only as poke_mode_only from airflow.sdk.bases.skipmixin import SkipMixin as SkipMixin from airflow.sdk.configuration import conf as conf + from airflow.sdk.definitions._internal.types import SET_DURING_EXECUTION as SET_DURING_EXECUTION from airflow.sdk.definitions.context import context_merge as context_merge from airflow.sdk.definitions.mappedoperator import MappedOperator as MappedOperator from airflow.sdk.definitions.template import literal as literal @@ -264,12 +265,23 @@ # ============================================================================ "Context": ("airflow.sdk", "airflow.utils.context"), "context_merge": ("airflow.sdk.definitions.context", "airflow.utils.context"), + # Default for a decorated operator's argument that the callable's return value fills in + "SET_DURING_EXECUTION": ( + "airflow.sdk.definitions._internal.types", + "airflow.providers.common.compat._set_during_execution", + ), "context_to_airflow_vars": ("airflow.sdk.execution_time.context", "airflow.utils.operator_helpers"), "AIRFLOW_VAR_NAME_FORMAT_MAPPING": ( "airflow.sdk.execution_time.context", "airflow.utils.operator_helpers", ), - "get_current_context": ("airflow.sdk", "airflow.operators.python"), + # On Airflow 2 the standard provider's version comes before core's: it raises RuntimeError + # outside a task, as Airflow 3 does, where core's raises AirflowException. + "get_current_context": ( + "airflow.sdk", + "airflow.providers.standard.operators.python", + "airflow.operators.python", + ), "get_parsing_context": ("airflow.sdk", "airflow.utils.dag_parsing_context"), # ============================================================================ # Timeout Utilities diff --git a/providers/common/compat/tests/unit/common/compat/test_sdk.py b/providers/common/compat/tests/unit/common/compat/test_sdk.py index d6f819ec9e5b0..8f0dd12985a31 100644 --- a/providers/common/compat/tests/unit/common/compat/test_sdk.py +++ b/providers/common/compat/tests/unit/common/compat/test_sdk.py @@ -22,6 +22,8 @@ import pytest +from airflow.providers.common.compat import sdk + from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS @@ -32,8 +34,6 @@ def test_all_compat_imports_work(): For each item, validates that at least one of the specified import paths works, ensuring the fallback mechanism is functional. """ - from airflow.providers.common.compat import sdk - failed_imports = [] for name in sdk.__all__: @@ -51,11 +51,9 @@ def test_all_compat_imports_work(): @pytest.mark.skipif(AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow < 3.0") -@pytest.mark.parametrize("name", ["BaseBranchOperator", "BranchMixIn"]) -def test_branching_imports_work_without_standard_provider(name, monkeypatch): +@pytest.mark.parametrize("name", ["BaseBranchOperator", "BranchMixIn", "get_current_context"]) +def test_airflow2_fallbacks_work_without_standard_provider(name, monkeypatch): """On Airflow 2 the standard provider is optional, so core paths must be used as fallback.""" - from airflow.providers.common.compat import sdk - real_import = builtins.__import__ def fake_import(module_name, *args, **kwargs): @@ -68,9 +66,22 @@ def fake_import(module_name, *args, **kwargs): assert getattr(sdk, name) is not None +def test_get_current_context_outside_a_task_raises_runtime_error(): + """With the standard provider installed, Airflow 2 matches Airflow 3; core's version raises AirflowException.""" + with pytest.raises(RuntimeError, match="no context was found"): + sdk.get_current_context() + + +@pytest.mark.skipif(AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow < 3.0") +def test_set_during_execution_renders_without_an_object_address(): + """Airflow 2 stores an unencodable template field as ``str(value)``, which must not vary per process.""" + from airflow.utils.types import ArgNotSet # Airflow 2 only; a deprecated redirect on Airflow 3 + + assert isinstance(sdk.SET_DURING_EXECUTION, ArgNotSet) + assert str(sdk.SET_DURING_EXECUTION) == "DYNAMIC (set during execution)" + + def test_invalid_import_raises_attribute_error(): """Test that importing non-existent attribute raises AttributeError.""" - from airflow.providers.common.compat import sdk - with pytest.raises(AttributeError, match="has no attribute 'NonExistentClass'"): _ = sdk.NonExistentClass diff --git a/uv.lock b/uv.lock index eaa61ed383d8e..55309aaca50e5 100644 --- a/uv.lock +++ b/uv.lock @@ -4642,6 +4642,7 @@ dependencies = [ { name = "apache-airflow-providers-common-compat" }, { name = "apache-airflow-providers-standard" }, { name = "pydantic-ai-slim" }, + { name = "structlog" }, ] [package.optional-dependencies] @@ -4771,6 +4772,7 @@ requires-dist = [ { name = "pypdf", marker = "extra == 'pdf'", specifier = ">=4.0.0" }, { name = "python-docx", marker = "extra == 'docx'", specifier = ">=1.0.0" }, { name = "sqlglot", marker = "extra == 'sql'", specifier = ">=30.0.0" }, + { name = "structlog", specifier = ">=24.2.0" }, { name = "typesafe-sdk", marker = "extra == 'typesafe'", specifier = ">=0.6.0" }, ] provides-extras = ["anthropic", "bedrock", "google", "openai", "typesafe", "mcp", "modal", "opensandbox", "code-mode", "shields", "skills", "avro", "parquet", "sql", "common-sql", "langchain", "llamaindex", "pdf", "docx", "git"] From 6144f1edbf2d0222e3b84bc542f26025f0f0de1d Mon Sep 17 00:00:00 2001 From: Kaxil Naik Date: Thu, 1 Oct 2026 00:20:40 +0100 Subject: [PATCH 2/3] Add tests for the Airflow 2 SET_DURING_EXECUTION stand-in --- .../compat/test__set_during_execution.py | 43 +++++++++++++++++++ .../tests/unit/common/compat/test_sdk.py | 9 ---- 2 files changed, 43 insertions(+), 9 deletions(-) create mode 100644 providers/common/compat/tests/unit/common/compat/test__set_during_execution.py diff --git a/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py b/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py new file mode 100644 index 0000000000000..077e8f6068f67 --- /dev/null +++ b/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py @@ -0,0 +1,43 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import pytest + +from airflow.providers.common.compat import sdk + +from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS + +pytestmark = pytest.mark.skipif(AIRFLOW_V_3_0_PLUS, reason="The stand-in is only used on Airflow 2") + +if not AIRFLOW_V_3_0_PLUS: + from airflow.providers.common.compat._set_during_execution import SET_DURING_EXECUTION + from airflow.serialization.helpers import serialize_template_field + from airflow.utils.types import ArgNotSet + + +def test_compat_sdk_hands_out_the_stand_in(): + assert sdk.SET_DURING_EXECUTION is SET_DURING_EXECUTION + + +def test_is_an_arg_not_set_sentinel(): + assert isinstance(SET_DURING_EXECUTION, ArgNotSet) + + +def test_serializes_as_the_airflow_3_sentinel_does(): + """A bare ``NOTSET`` serializes as an object address, which differs in every process.""" + assert serialize_template_field(SET_DURING_EXECUTION, "prompt") == "DYNAMIC (set during execution)" diff --git a/providers/common/compat/tests/unit/common/compat/test_sdk.py b/providers/common/compat/tests/unit/common/compat/test_sdk.py index 8f0dd12985a31..adc72aea2de59 100644 --- a/providers/common/compat/tests/unit/common/compat/test_sdk.py +++ b/providers/common/compat/tests/unit/common/compat/test_sdk.py @@ -72,15 +72,6 @@ def test_get_current_context_outside_a_task_raises_runtime_error(): sdk.get_current_context() -@pytest.mark.skipif(AIRFLOW_V_3_0_PLUS, reason="Test requires Airflow < 3.0") -def test_set_during_execution_renders_without_an_object_address(): - """Airflow 2 stores an unencodable template field as ``str(value)``, which must not vary per process.""" - from airflow.utils.types import ArgNotSet # Airflow 2 only; a deprecated redirect on Airflow 3 - - assert isinstance(sdk.SET_DURING_EXECUTION, ArgNotSet) - assert str(sdk.SET_DURING_EXECUTION) == "DYNAMIC (set during execution)" - - def test_invalid_import_raises_attribute_error(): """Test that importing non-existent attribute raises AttributeError.""" with pytest.raises(AttributeError, match="has no attribute 'NonExistentClass'"): From de799095c4458a4d0f20164342f186dd3eeac67f Mon Sep 17 00:00:00 2001 From: Kaxil Naik Date: Thu, 1 Oct 2026 09:46:42 +0100 Subject: [PATCH 3/3] Address review: rename run-key helper, scope docs heading to 2.11 - Rename task_instance_run_key to make_task_instance_run_key. - Title the installation section "Airflow 2.11", so it does not read as covering every Airflow 2 release. - Silence the Airflow 2-only ArgNotSet import for mypy on Airflow 3. --- providers/common/ai/docs/installation.rst | 4 ++-- .../ai/src/airflow/providers/common/ai/observability.py | 2 +- .../ai/src/airflow/providers/common/ai/operators/agent.py | 6 +++--- .../common/ai/tests/unit/common/ai/test_observability.py | 4 ++-- .../tests/unit/common/compat/test__set_during_execution.py | 2 +- 5 files changed, 9 insertions(+), 9 deletions(-) diff --git a/providers/common/ai/docs/installation.rst b/providers/common/ai/docs/installation.rst index 1f8c61086131e..483defafdf683 100644 --- a/providers/common/ai/docs/installation.rst +++ b/providers/common/ai/docs/installation.rst @@ -89,8 +89,8 @@ The provider runs on Airflow 2.11, but some features need a newer core: Pydantic model; on older cores it arrives as a ``dict`` - Airflow 3.3 -Airflow 2 ---------- +Airflow 2.11 +------------ On Airflow 2.11 the operators, decorators, hooks and toolsets run as they do on Airflow 3.0, apart from the table above. Three things differ from an Airflow 3 install: diff --git a/providers/common/ai/src/airflow/providers/common/ai/observability.py b/providers/common/ai/src/airflow/providers/common/ai/observability.py index 0d70687b28e43..e28d453af5b1a 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/observability.py +++ b/providers/common/ai/src/airflow/providers/common/ai/observability.py @@ -135,7 +135,7 @@ def build_run_identity_attributes(ti: Any) -> dict[str, Any]: return attributes -def task_instance_run_key(ti: Any) -> str: +def make_task_instance_run_key(ti: Any) -> str: """ Return a per-attempt key for ``ti``: its id on Airflow 3, a composite on Airflow 2. diff --git a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py index 09e464844e939..4f13ad4bf0d8a 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py @@ -47,8 +47,8 @@ from airflow.providers.common.ai.mixins.hitl_review import HITLReviewMixin from airflow.providers.common.ai.observability import ( build_run_identity_attributes, + make_task_instance_run_key, stamp_identity_on_agent_spans, - task_instance_run_key, ) from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset from airflow.providers.common.ai.utils.logging import ( @@ -999,7 +999,7 @@ def _report_failed_run(self, context: Context, run_usage: RunUsage) -> None: return ti = context["task_instance"] try: - ti.xcom_push(key="run_id", value=task_instance_run_key(ti)) + ti.xcom_push(key="run_id", value=make_task_instance_run_key(ti)) except Exception: self.log.warning("Failed to push run_id XCom for the failed run", exc_info=True) if attempt_usage is not None: @@ -1113,7 +1113,7 @@ def execute(self, context: Context) -> Any: # each retry; dag/run/task/map/try on Airflow 2) is a unique, reverse-resolvable # join key. It lands on result.run_id, the run's messages, and the # ``gen_ai.agent.call.id`` span attribute. - run_kwargs: dict[str, Any] = {"usage_limits": usage_limits, "run_id": task_instance_run_key(ti)} + run_kwargs: dict[str, Any] = {"usage_limits": usage_limits, "run_id": make_task_instance_run_key(ti)} history = self._resolve_message_history() if history is not None: run_kwargs["message_history"] = history diff --git a/providers/common/ai/tests/unit/common/ai/test_observability.py b/providers/common/ai/tests/unit/common/ai/test_observability.py index 6ba883e9b6d49..946fb74e8aad0 100644 --- a/providers/common/ai/tests/unit/common/ai/test_observability.py +++ b/providers/common/ai/tests/unit/common/ai/test_observability.py @@ -173,13 +173,13 @@ class TestTaskInstanceRunKey: def test_uses_the_task_instance_id_when_it_has_one(self): ti = SimpleNamespace(id="0199-uuid", dag_id="d", task_id="t", run_id="r", try_number=2, map_index=3) - assert observability.task_instance_run_key(ti) == "0199-uuid" + assert observability.make_task_instance_run_key(ti) == "0199-uuid" def test_builds_a_per_attempt_key_without_an_id(self): """Airflow 2 task instances have no ``id`` column.""" ti = SimpleNamespace(dag_id="d", task_id="t", run_id="r", try_number=2, map_index=3) - assert observability.task_instance_run_key(ti) == "d/r/t/3/2" + assert observability.make_task_instance_run_key(ti) == "d/r/t/3/2" class TestStampIdentityOnAgentSpans: diff --git a/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py b/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py index 077e8f6068f67..a50fa1813e8a2 100644 --- a/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py +++ b/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py @@ -27,7 +27,7 @@ if not AIRFLOW_V_3_0_PLUS: from airflow.providers.common.compat._set_during_execution import SET_DURING_EXECUTION from airflow.serialization.helpers import serialize_template_field - from airflow.utils.types import ArgNotSet + from airflow.utils.types import ArgNotSet # type: ignore[attr-defined] # Airflow 2 only def test_compat_sdk_hands_out_the_stand_in():