This is an automated email from the ASF dual-hosted git repository.
kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 8f3e8466c67 Support Airflow 2.11 in the Common AI provider (#73991)
8f3e8466c67 is described below
commit 8f3e8466c67c0d4bf5db4bdbb99fdce50f436e4f
Author: Kaxil Naik <[email protected]>
AuthorDate: Thu Oct 1 15:10:39 2026 +0100
Support Airflow 2.11 in the Common AI provider (#73991)
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.
---
dev/breeze/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 +
providers/common/ai/docs/operators/llm_batch.rst | 4 +-
providers/common/ai/docs/quickstart.rst | 3 +-
providers/common/ai/docs/self_hosted_models.rst | 3 +-
providers/common/ai/pyproject.toml | 7 +-
.../ai/src/airflow/providers/common/ai/__init__.py | 4 +-
.../airflow/providers/common/ai/batch/anthropic.py | 5 +-
.../airflow/providers/common/ai/batch/openai.py | 5 +-
.../airflow/providers/common/ai/batch/results.py | 5 +-
.../providers/common/ai/decorators/agent.py | 2 +-
.../airflow/providers/common/ai/decorators/llm.py | 2 +-
.../providers/common/ai/decorators/llm_batch.py | 2 +-
.../providers/common/ai/decorators/llm_branch.py | 2 +-
.../common/ai/decorators/llm_file_analysis.py | 2 +-
.../common/ai/decorators/llm_schema_compare.py | 2 +-
.../providers/common/ai/decorators/llm_sql.py | 2 +-
.../providers/common/ai/durable/caching_model.py | 4 +-
.../providers/common/ai/durable/caching_toolset.py | 4 +-
.../providers/common/ai/durable/fingerprint.py | 4 +-
.../airflow/providers/common/ai/durable/storage.py | 7 +-
.../common/ai/durable/task_state_store.py | 4 +-
.../providers/common/ai/hooks/pydantic_ai.py | 10 +++
.../airflow/providers/common/ai/observability.py | 19 ++++-
.../airflow/providers/common/ai/operators/agent.py | 16 +++--
.../common/ai/operators/llamaindex_embedding.py | 2 +-
.../common/ai/operators/llamaindex_retrieval.py | 2 +-
.../providers/common/ai/utils/task_logger.py | 67 ++++++++++++++++++
.../providers/common/ai/utils/usage_budget.py | 5 +-
.../ai/tests/unit/common/ai/batch/test_dispatch.py | 2 +-
.../ai/tests/unit/common/ai/batch/test_results.py | 2 +-
.../ai/tests/unit/common/ai/batch/test_state.py | 2 +-
.../unit/common/ai/decorators/test_llm_batch.py | 3 +-
.../unit/common/ai/durable/test_replay_cost.py | 2 +-
.../common/ai/durable/test_replay_verification.py | 2 +-
.../tests/unit/common/ai/durable/test_storage.py | 2 +-
.../tests/unit/common/ai/hooks/test_pydantic_ai.py | 15 ++++
.../tests/unit/common/ai/mixins/test_approval.py | 2 +-
.../tests/unit/common/ai/operators/test_agent.py | 32 +++++++--
.../common/ai/operators/test_document_loader.py | 12 ++--
.../ai/operators/test_llamaindex_embedding.py | 2 +-
.../ai/operators/test_llamaindex_retrieval.py | 4 +-
.../unit/common/ai/operators/test_llm_batch.py | 3 +-
.../ai/tests/unit/common/ai/test_observability.py | 19 +++++
.../unit/common/ai/toolsets/test_object_storage.py | 8 ++-
.../tests/unit/common/ai/utils/test_task_logger.py | 82 ++++++++++++++++++++++
.../common/compat/_set_during_execution.py | 40 +++++++++++
.../src/airflow/providers/common/compat/sdk.py | 14 +++-
.../common/compat/test__set_during_execution.py | 43 ++++++++++++
.../compat/tests/unit/common/compat/test_sdk.py | 18 ++---
uv.lock | 2 +
54 files changed, 463 insertions(+), 88 deletions(-)
diff --git a/dev/breeze/src/airflow_breeze/global_constants.py
b/dev/breeze/src/airflow_breeze/global_constants.py
index 43dfcf7e7ed..6befa6a0566 100644
--- a/dev/breeze/src/airflow_breeze/global_constants.py
+++ b/dev/breeze/src/airflow_breeze/global_constants.py
@@ -862,7 +862,7 @@ PROVIDERS_COMPATIBILITY_TESTS_MATRIX: list[dict[str, str |
list[str]]] = [
{
"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 9d196fafb8d..b1204abc6c8 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 666450bafcc..fb2d2d62caa 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 43e12ca82c8..483defafdf6 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 <approval_gates>` and :doc:`HITL review
<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 <retry_policies>`
- Airflow 3.3
* - :doc:`Durable execution <durable_execution>` without configuring
``[common.ai] durable_cache_path`` (the task state store)
- Airflow 3.3
+ * - :doc:`Tool approval <tool_approval>` that pauses the task; on older
cores a tool
+ marked for approval fails the task
+ - Airflow 3.3
+ * - A :doc:`structured output <structured_output>` reaching downstream
tasks as the
+ Pydantic model; on older cores it arrives as a ``dict``
+ - Airflow 3.3
+
+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:
+
+* 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 828b28d1dda..3a79c2c64b1 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
``<task-instance id>-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
+ ``<dag_id>/<run_id>/<task_id>/<map_index>/<try_number>``, 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 806517c4051..c9030bf7297 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 a21393ed190..c32379cc4e0 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
<apache-airflow:installation/index>` 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 5eda43b7a3b..ef1244f78f2 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 dd818c99ad9..18d50b3033e 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 db65dd59825..a9ebbadf0fd 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 @@ __all__ = ["__version__"]
__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 1fc58b45aeb..d928026dddb 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 @@ import json
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 @@ from airflow.providers.common.ai.batch.base import (
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 11273789c3a..323ff9062f5 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 collections.abc import Iterator
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 @@ from airflow.providers.common.ai.batch.base import (
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 97e771d6567..3afba5d23bb 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 collections.abc import Iterable, Mapping
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 25272d06009..d2a92e07668 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 @@ from airflow.providers.common.ai.utils.validation import (
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 8f29226202a..5e3ec84c6d1 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 @@ from airflow.providers.common.ai.utils.validation import (
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 d7f351b3728..d1bc139a1b4 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 typing import TYPE_CHECKING, Any, ClassVar
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 d835feeab17..fd8ee3f2727 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 @@ from airflow.providers.common.ai.utils.validation import (
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 305184fb0d1..83447e92322 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 typing import TYPE_CHECKING, Any, ClassVar
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 2e106509215..a4f22bfa486 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 @@ from airflow.providers.common.ai.utils.validation import (
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 7ea9765ddcd..e06a46183b8 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 @@ from airflow.providers.common.ai.utils.validation import (
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 ef1d6940f9c..07084d5fdbd 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 __future__ import annotations
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 fa4a61b9af7..1eb30127c01 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 __future__ import annotations
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 @@ if TYPE_CHECKING:
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 bc4a291ea92..be3c2b6f1c2 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 hashlib
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 a3f974bcb42..4e98fe40eb4 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 @@ import json
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 ecf1c0e066d..c02a33881c1 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 @@ from __future__ import annotations
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 @@ if TYPE_CHECKING:
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 c8531d4fe12..efb34edaac0 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 @@ class PydanticAIHook(BaseHook):
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 25e4a8712b7..e28d453af5b 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 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.
+
+ 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 974c7697cdd..4f13ad4bf0d 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,6 +47,7 @@ from airflow.providers.common.ai.mixins.cancellable_run
import CancellableAgentR
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,
)
from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset
@@ -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 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator,
HITLReviewMixin):
return
ti = context["task_instance"]
try:
- ti.xcom_push(key="run_id", value=str(ti.id))
+ 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:
@@ -1106,10 +1109,11 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
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":
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/src/airflow/providers/common/ai/operators/llamaindex_embedding.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py
index c7003db1e33..8a878a874f7 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 @@ class LlamaIndexEmbeddingOperator(BaseOperator):
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 234a3bfdd4d..d61e1c5f2be 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 @@ class LlamaIndexRetrievalOperator(BaseOperator):
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 00000000000..5bb9c07bef2
--- /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 9e2f1fb6067..f67c686fa99 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 @@ import dataclasses
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 430fa49ad46..1603f6b2c05 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 @@ from airflow.providers.common.ai.exceptions import (
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 82da64ac3a5..1cbb218a5a4 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 @@ from airflow.providers.common.ai.batch.results import (
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 fb91d5f5fab..7b157f1e907 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 @@ from airflow.providers.common.ai.batch.state import (
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 538148c8faf..f66bb87285e 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.batch.base import (
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 3b083a2a0a1..4905d44c22a 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.caching_model import
CachingModel
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 ad9c244b0c6..1e57ad45eaa 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_model import
CachingModel
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 c1a4176d164..c2670655195 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.messages import (
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 2a4dc2b1fec..5f856c1d5c0 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 960999f0495..85123748ff6 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 @@ from airflow.providers.common.ai.mixins.approval import (
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 be384d21a05..e344947308e 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 @@ from airflow.providers.common.ai.utils.usage_budget import (
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 @@ class _InMemoryDurableStorage:
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 @@ class TestAgentOperatorToolsetTemplating:
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 @@ class TestAgentOperatorExecute:
"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 60031e3edfb..dcd0ab8dac9 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 @@ from unittest.mock import MagicMock, patch
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 @@ class TestFileDiscovery:
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 @@ class TestCloudUriDispatch:
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 @@ class TestCloudUriDispatch:
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 @@ class TestCloudUriDispatch:
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 @@ class TestCloudUriDispatch:
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 e410dfded05..0d647f372e0 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 @@ class TestEmbeddingOperatorPersist:
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
``<scheme>://<conn_id>@<bucket>/...``
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 c0c4b0c9e89..6c1378c5c6f 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 @@ class TestRetrievalOperatorMissingIndex:
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 @@ class TestRetrievalOperatorMissingIndex:
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
``<scheme>://<conn_id>@<bucket>/...``
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 f8703b9725b..fbe2a041819 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.exceptions import (
)
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 91f4e8c91df..946fb74e8aa 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 @@ class TestBuildRunIdentityAttributes:
"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.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.make_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 989552db68f..0dc15cdd5b4 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.models.test import TestModel
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 00000000000..f48b0b6a7cd
--- /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)
+
+
[email protected]
+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 00000000000..729fc3db164
--- /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 de7480c7336..f37a194b337 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 @@ if TYPE_CHECKING:
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 @@ _IMPORT_MAP: dict[str, str | tuple[str, ...]] = {
#
============================================================================
"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__set_during_execution.py
b/providers/common/compat/tests/unit/common/compat/test__set_during_execution.py
new file mode 100644
index 00000000000..a50fa1813e8
--- /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 # type: ignore[attr-defined] #
Airflow 2 only
+
+
+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 d6f819ec9e5..adc72aea2de 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 builtins
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")
[email protected]("name", ["BaseBranchOperator", "BranchMixIn"])
-def test_branching_imports_work_without_standard_provider(name, monkeypatch):
[email protected]("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,13 @@ def
test_branching_imports_work_without_standard_provider(name, monkeypatch):
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()
+
+
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 eaa61ed383d..55309aaca50 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"]