This is an automated email from the ASF dual-hosted git repository.
Lee-W 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 1050685b721 Add connection-driven provider failover for common.ai LLM
calls (#72156)
1050685b721 is described below
commit 1050685b72110a6dd25dd56e545f6bbbf02b1946
Author: Wei Lee <[email protected]>
AuthorDate: Sat Sep 19 16:12:46 2026 +0900
Add connection-driven provider failover for common.ai LLM calls (#72156)
---
providers/common/ai/docs/changelog.rst | 11 +
.../common/ai/docs/connections/pydantic_ai.rst | 20 +
.../ai/docs/connections/pydantic_ai_azure.rst | 24 +-
.../ai/docs/connections/pydantic_ai_bedrock.rst | 14 +
.../ai/docs/connections/pydantic_ai_vertex.rst | 24 +-
providers/common/ai/docs/examples.rst | 4 +
providers/common/ai/docs/index.rst | 1 +
providers/common/ai/docs/provider_fallback.rst | 231 ++++++
providers/common/ai/docs/retry_policies.rst | 34 +
providers/common/ai/provider.yaml | 36 +
.../common/ai/example_dags/example_llm_fallback.py | 101 +++
.../providers/common/ai/get_provider_info.py | 22 +-
.../providers/common/ai/hooks/pydantic_ai.py | 403 ++++++++-
.../airflow/providers/common/ai/operators/agent.py | 11 +
.../airflow/providers/common/ai/operators/llm.py | 11 +
.../providers/common/ai/operators/llm_branch.py | 7 +
.../common/ai/operators/llm_file_analysis.py | 7 +
.../common/ai/operators/llm_schema_compare.py | 7 +
.../providers/common/ai/operators/llm_sql.py | 7 +
.../ai/tests/unit/common/ai/decorators/test_llm.py | 19 +
.../tests/unit/common/ai/hooks/test_pydantic_ai.py | 897 ++++++++++++++++++++-
.../tests/unit/common/ai/operators/test_agent.py | 47 +-
.../ai/tests/unit/common/ai/operators/test_llm.py | 56 +-
.../common/ai/operators/test_llm_file_analysis.py | 1 +
.../tests/unit/common/ai/operators/test_llm_sql.py | 20 +
.../ai/tests/unit/common/ai/utils/test_logging.py | 29 +
26 files changed, 1994 insertions(+), 50 deletions(-)
diff --git a/providers/common/ai/docs/changelog.rst
b/providers/common/ai/docs/changelog.rst
index b8bc72647cf..1eb0b82aca7 100644
--- a/providers/common/ai/docs/changelog.rst
+++ b/providers/common/ai/docs/changelog.rst
@@ -25,6 +25,17 @@
Changelog
---------
+.. note::
+ Configuring ``fallback_conn_ids`` on a connection (or the matching
operator/decorator
+ argument) changes what exception a task raises once every connection in the
chain fails:
+ it is ``pydantic_ai.exceptions.FallbackExceptionGroup``, not the last
provider's own
+ exception. An existing ``RetryRule(exception=ModelHTTPError, ...)`` -- in
+ ``LLMRetryPolicy.fallback_rules`` or a plain ``ExceptionRetryPolicy`` --
stops matching
+ as soon as the connection gains a fallback chain, with no change to the Dag
needed to
+ trigger it. Match ``pydantic_ai.exceptions.FallbackExceptionGroup``
explicitly as well,
+ or inspect its ``.exceptions`` attribute for the original per-model errors.
See
+ :doc:`retry_policies`, "When the connection also carries a fallback chain".
+
0.9.0
.....
diff --git a/providers/common/ai/docs/connections/pydantic_ai.rst
b/providers/common/ai/docs/connections/pydantic_ai.rst
index 19a48988da9..f82f5d4042b 100644
--- a/providers/common/ai/docs/connections/pydantic_ai.rst
+++ b/providers/common/ai/docs/connections/pydantic_ai.rst
@@ -38,6 +38,11 @@ Model
dedicated input in the connection form (via ``conn-fields``) and stores its
value in ``extra["model"]``.
+ The ``provider:`` prefix is required here: this generic connection type has
+ no platform of its own (unlike the vendor connection types below), so a
+ bare name (e.g. ``gpt-5.6-sol`` without ``openai:``) raises ``ValueError``
+ naming this connection rather than being resolved automatically.
+
Examples: ``openai:gpt-5.6-sol``, ``anthropic:claude-sonnet-5``,
``bedrock:us.anthropic.claude-opus-4-6-v1:0``, ``google:gemini-2.0-flash``
@@ -76,6 +81,12 @@ Extra (JSON, optional)
When using the UI, the "Model" field above writes to this same location
automatically.
+Fallback Connections
+ Other connection IDs to fail over to, in order, while this provider is
+ unavailable. Stored in ``extra["fallback_conn_ids"]``. Entries may name any
+ ``pydanticai`` connection type, so one chain can span vendors. See
+ :doc:`/provider_fallback`.
+
Examples
--------
@@ -153,3 +164,12 @@ The hook reads the model from these sources in priority
order:
1. ``model_id`` parameter on the hook/operator
2. ``model`` in the connection's extra JSON (set by the "Model" conn-field in
the UI)
+3. When this connection is used as a fallback and neither of the above is set,
the
+ *bare* ``model_id`` forwarded from the primary connection (see
:doc:`/provider_fallback`) --
+ a forwarded name that already pins a platform is not applied here, since it
names a
+ model of the primary's own platform.
+
+Whichever name is chosen, a name that already pins a recognized platform (its
segment
+before the first ``:`` is itself a pydantic-ai provider) is used verbatim; a
bare name is
+qualified with this connection's platform, and this generic connection type
has none,
+so a bare name reaching this step always raises.
diff --git a/providers/common/ai/docs/connections/pydantic_ai_azure.rst
b/providers/common/ai/docs/connections/pydantic_ai_azure.rst
index 8f7b9167060..692bb5909ba 100644
--- a/providers/common/ai/docs/connections/pydantic_ai_azure.rst
+++ b/providers/common/ai/docs/connections/pydantic_ai_azure.rst
@@ -61,12 +61,18 @@ Configuring the Connection
--------------------------
Model
- Azure model identifier (e.g. ``azure:gpt-4o``). This field appears as a
- dedicated input in the connection form (via ``conn-fields``) and stores its
- value in ``extra["model"]``.
-
- The ``azure:`` prefix is required — it is what makes pydantic-ai
instantiate
- the Azure OpenAI provider instead of the plain OpenAI one.
+ Azure model identifier (e.g. ``azure:gpt-4o``, or the bare ``gpt-4o``).
This
+ field appears as a dedicated input in the connection form (via
+ ``conn-fields``) and stores its value in ``extra["model"]``.
+
+ A bare name is automatically resolved to ``azure:<name>`` -- Azure OpenAI
is
+ this connection type's own platform, so nothing else needs naming it
+ explicitly. Writing the ``azure:`` prefix yourself has the same effect and
is
+ still accepted. A name prefixed with a *different*, recognized platform
(e.g.
+ ``openai:gpt-4o``) is used verbatim instead, pinning that platform and
+ bypassing Azure OpenAI entirely -- a name is only treated as already
prefixed
+ when the segment before its first ``:`` is itself a real pydantic-ai
+ provider, not merely present.
API Key (Password field)
The Azure OpenAI API key.
@@ -82,6 +88,12 @@ API Version (Extra field)
``OPENAI_API_VERSION`` environment variable if omitted. Endpoints matching
either OpenAI-compatible v1 form reject this field.
+Fallback Connections
+ Other connection IDs to fail over to, in order, while this provider is
+ unavailable. Stored in ``extra["fallback_conn_ids"]``. Entries may name any
+ ``pydanticai`` connection type, so one chain can span vendors. See
+ :doc:`/provider_fallback`.
+
Examples
--------
diff --git a/providers/common/ai/docs/connections/pydantic_ai_bedrock.rst
b/providers/common/ai/docs/connections/pydantic_ai_bedrock.rst
index 4815bf6d389..4c106b8490c 100644
--- a/providers/common/ai/docs/connections/pydantic_ai_bedrock.rst
+++ b/providers/common/ai/docs/connections/pydantic_ai_bedrock.rst
@@ -64,6 +64,14 @@ All fields below are ``extra`` (JSON) fields.
Model
Bedrock model identifier (e.g. ``bedrock:us.anthropic.claude-opus-4-5``).
+ A bare name is automatically resolved to ``bedrock:<name>`` -- Bedrock is
this
+ connection type's own platform. This includes Bedrock's version-suffixed
ids,
+ which contain a ``:`` of their own (e.g.
``us.anthropic.claude-opus-4-6-v1:0``):
+ that ``:`` is not a recognized pydantic-ai provider name, so it does not
count
+ as an existing platform prefix, and the whole bare id still gets
``bedrock:``
+ prepended (``bedrock:us.anthropic.claude-opus-4-6-v1:0``). Writing the
+ ``bedrock:`` prefix yourself has the same effect and is still accepted.
+
AWS Region
AWS region (e.g. ``us-east-1``). Falls back to the ``AWS_DEFAULT_REGION``
environment variable.
@@ -93,6 +101,12 @@ Read Timeout (s)
Connect Timeout (s)
boto3 connect timeout in seconds (float, optional).
+Fallback Connections
+ Other connection IDs to fail over to, in order, while this provider is
+ unavailable. Stored in ``extra["fallback_conn_ids"]``. Entries may name any
+ ``pydanticai`` connection type, so one chain can span vendors. See
+ :doc:`/provider_fallback`.
+
Credentials
-----------
diff --git a/providers/common/ai/docs/connections/pydantic_ai_vertex.rst
b/providers/common/ai/docs/connections/pydantic_ai_vertex.rst
index 5d91bb4ec7b..03503960491 100644
--- a/providers/common/ai/docs/connections/pydantic_ai_vertex.rst
+++ b/providers/common/ai/docs/connections/pydantic_ai_vertex.rst
@@ -63,11 +63,19 @@ Configuring the Connection
All fields below are ``extra`` (JSON) fields.
Model
- Google model identifier (e.g. ``google-cloud:gemini-2.0-flash``). The
- ``google-cloud:`` prefix is required — it is what makes pydantic-ai
- instantiate the ``GoogleCloudProvider``, which is what accepts this
- hook's ``project`` / ``location`` / ``service_account_info`` fields (see
- "Credentials" below).
+ Google model identifier (e.g. ``google-cloud:gemini-2.0-flash``, or the
+ bare ``gemini-2.0-flash``). A bare name is automatically resolved to
+ ``google-cloud:<name>``, instantiating the ``GoogleCloudProvider`` that
+ accepts this hook's ``project`` / ``location`` / ``service_account_info``
+ fields (see "Credentials" below) -- Vertex AI is this connection type's
+ default platform for a bare name, and that default holds regardless of
+ which credential fields are set on the connection: it is **not** inferred
+ from whether ``api_key`` is present, because ``api_key`` here can equally
+ mean Vertex AI Express Mode credentials (see "Credentials" below), so its
+ presence alone cannot tell the two platforms apart. To reach the
+ Generative Language API instead, prefix the model explicitly with
+ ``google:`` -- that spelling routes to a different provider regardless of
+ which fields this connection sets.
GCP Project
Google Cloud project ID. Falls back to the ``GOOGLE_CLOUD_PROJECT``
@@ -108,6 +116,12 @@ Service Account Info
Custom Endpoint URL
Override the Google API base URL (optional).
+Fallback Connections
+ Other connection IDs to fail over to, in order, while this provider is
+ unavailable. Stored in ``extra["fallback_conn_ids"]``. Entries may name any
+ ``pydanticai`` connection type, so one chain can span vendors. See
+ :doc:`/provider_fallback`.
+
Credentials
-----------
diff --git a/providers/common/ai/docs/examples.rst
b/providers/common/ai/docs/examples.rst
index a74550d5e60..536e3f1b46a 100644
--- a/providers/common/ai/docs/examples.rst
+++ b/providers/common/ai/docs/examples.rst
@@ -140,6 +140,10 @@ Reliability
* - :doc:`retry_policies`
- Classifying task failures with an LLM to decide retry, fail, or delay.
Source:
`example_llm_retry_policy.py
<https://github.com/apache/airflow/blob/providers-common-ai/|version|/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_retry_policy.py>`__.
+ * - :doc:`provider_fallback`
+ - Failing over to another vendor inside one task attempt, and drilling
the chain
+ without waiting for an outage. Source:
+ `example_llm_fallback.py
<https://github.com/apache/airflow/blob/providers-common-ai/|version|/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_fallback.py>`__.
.. toctree::
:hidden:
diff --git a/providers/common/ai/docs/index.rst
b/providers/common/ai/docs/index.rst
index ff9b39d47ab..5526994cb48 100644
--- a/providers/common/ai/docs/index.rst
+++ b/providers/common/ai/docs/index.rst
@@ -159,6 +159,7 @@ See the Optional dependencies table below for the exact
package each extra insta
Toolsets <toolsets>
Operators <operators/index>
Examples <examples>
+ Provider fallback <provider_fallback>
Retry Policies <retry_policies>
Self-hosted models <self_hosted_models>
HITL Review <hitl_review>
diff --git a/providers/common/ai/docs/provider_fallback.rst
b/providers/common/ai/docs/provider_fallback.rst
new file mode 100644
index 00000000000..0d957b48a19
--- /dev/null
+++ b/providers/common/ai/docs/provider_fallback.rst
@@ -0,0 +1,231 @@
+ .. 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.
+
+Provider fallback
+=================
+
+A single ``llm_conn_id`` gives a task one provider. When that provider is
down, the task
+fails and retries into the same outage. ``fallback_conn_ids`` gives the
connection an
+ordered list of other connections to try, so a provider outage moves to the
next vendor
+inside the same task attempt.
+
+Configure it on the connection
+------------------------------
+
+Put the chain in the primary connection's extra:
+
+.. code-block:: json
+
+ {
+ "model": "openai:gpt-5",
+ "fallback_conn_ids": ["anthropic_prod", "bedrock_dr"]
+ }
+
+Every entry is an Airflow connection ID, resolved through the hook registered
for its own
+connection type. A chain can therefore mix vendors whose credentials live in
different
+connection fields — ``pydanticai`` for OpenAI, ``pydanticai_bedrock`` for a
Bedrock
+standby — without the Dag knowing anything about either.
+
+That is the point of configuring it here rather than in Dag code: the Dag
keeps naming one
+connection, and whoever administers the connections owns the failover
topology. Changing a
+standby provider is a connection edit, not a Dag deployment.
+
+A *bare* model name (e.g. ``"gpt-5"`` rather than ``"openai:gpt-5"``) is
forwarded down
+the chain as a logical model name: each connection that has no ``model`` of
its own
+resolves that name against its own platform, so one bare name can reach a
primary and
+every fallback without repeating it per connection. It does not matter where
the primary's
+name comes from -- the ``Model`` field on its connection and a ``model_id`` on
the operator
+or hook are forwarded alike. A fallback with its own ``model`` in
+extra always uses that instead -- this is how a fallback pins a spelling the
forwarded
+name would not produce, such as Bedrock's region-prefixed ``us.anthropic.``
model ids. A
+name that already pins a platform (its segment before the first ``:`` is
itself a
+recognized provider, e.g. ``"openai:gpt-5"``) is *not* forwarded; a fallback
with no
+``model`` of its own still raises "no model specified" rather than trying a
prefixed name
+meant for a different provider.
+
+A bare name with no ``:`` of its own (e.g. ``"gpt-5"``) forwards to any
fallback
+regardless of platform, since nothing about the spelling is vendor-specific. A
bare name
+that itself contains a ``:`` -- a vendor's own native model id, such as
Bedrock's
+version-suffixed ``"us.anthropic.claude-opus-4-6-v1:0"`` -- only forwards to a
fallback on
+the *same* platform: that spelling is only meaningful on the vendor that
produced it, so a
+Bedrock primary's native id reaches a Bedrock fallback but not an Azure one.
For example,
+a Bedrock primary with ``fallback_conn_ids: ["bedrock_dr"]`` forwards
+``"us.anthropic.claude-opus-4-6-v1:0"`` to ``bedrock_dr`` unchanged; the same
primary with
+``fallback_conn_ids: ["azure_dr"]`` does not forward it to ``azure_dr`` at
all, and
+``azure_dr`` raises "no model specified" unless its own extra sets a
``model``. See
+:doc:`connections/pydantic_ai_azure`, :doc:`connections/pydantic_ai_bedrock`
and
+:doc:`connections/pydantic_ai_vertex` for how each vendor connection resolves
a bare name.
+
+Configure it on the operator
+-----------------------------
+
+``fallback_conn_ids`` is also a parameter on
+:class:`~airflow.providers.common.ai.operators.llm.LLMOperator`,
+:class:`~airflow.providers.common.ai.operators.agent.AgentOperator`, their
subclasses,
+and the matching ``@task.llm`` / ``@task.agent`` decorators -- mirroring
``model_id``,
+which is settable at the same two layers:
+
+.. exampleinclude::
/../../ai/src/airflow/providers/common/ai/example_dags/example_llm_fallback.py
+ :language: python
+ :dedent: 0
+ :start-after: [START howto_llm_fallback_operator_argument]
+ :end-before: [END howto_llm_fallback_operator_argument]
+
+The operator argument overrides the connection's extra field, and passing
``[]``
+explicitly disables a chain configured there -- ``None`` (the default) reads
whatever
+the connection says. Use this when a task should own its own failover order
instead of
+inheriting it from however the connection is configured.
+
+Configure it in code
+--------------------
+
+:class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook` also
takes the list
+directly, which is what a task that constructs the hook itself (rather than
through an
+operator) should use:
+
+.. exampleinclude::
/../../ai/src/airflow/providers/common/ai/example_dags/example_llm_fallback.py
+ :language: python
+ :dedent: 0
+ :start-after: [START howto_llm_fallback_hook_argument]
+ :end-before: [END howto_llm_fallback_hook_argument]
+
+The argument wins over the connection's extra, and passing ``[]`` explicitly
disables a
+chain configured there. Omitting it entirely (``None``) means "use whatever
the connection
+says", which is why the two are not interchangeable.
+
+Where this sits among the retry layers
+--------------------------------------
+
+Three mechanisms handle failure at different time scales, and they compose
rather than
+replace each other:
+
+.. list-table::
+ :header-rows: 1
+ :widths: 25 40 35
+
+ * - Scope
+ - Mechanism
+ - Handles
+ * - Within one model call
+ - ``fallback_conn_ids``
+ - This vendor's API is returning errors; ask the next one (any
``ModelAPIError``,
+ transient or not)
+ * - Within one task attempt
+ - ``timeout`` in pydantic-ai's ``ModelSettings``
+ - This vendor is slow rather than down
+ * - Across task attempts
+ - :doc:`retry_policies` (including ``LLMRetryPolicy``)
+ - Whether this failure is worth retrying at all
+
+A chain does not remove the need for the outer layers. It covers the case
where another
+vendor can answer the same prompt now; a bad prompt, an exhausted quota on
every vendor, or
+a permanent data error still has to be decided by the retry policy.
+
+Adding a chain changes what the retry layer sees. When every connection in the
chain
+fails, the exception the task raises is
``pydantic_ai.exceptions.FallbackExceptionGroup``,
+not the last provider's own exception, so retry rules matched against a
provider-specific
+exception type stop matching. Before adding a chain to a connection that Dags
already use,
+read :doc:`retry_policies` -- the section "When the connection also carries a
fallback
+chain" spells out what to check.
+
+Costs to know before configuring a long chain
+---------------------------------------------
+
+**The timeout multiplies.** pydantic-ai applies a ``ModelSettings`` timeout to
each model
+in the chain, not to the chain as a whole. A 30-second timeout across three
connections is
+a 90-second worst case for one call.
+
+**There is no circuit breaker.** Every call tries the primary first. During an
outage each
+task instance pays the primary's timeout again before failing over, so 500
mapped tasks pay
+it 500 times. Keeping the primary's timeout short bounds both of these.
+
+**Chains are not resolved recursively.** If a connection listed as a fallback
declares its
+own ``fallback_conn_ids``, resolution fails with an error rather than
following it. List
+every provider directly on the primary; a flat chain is the one you can read
off a single
+connection.
+
+**A malformed prompt walks the whole chain.** Failover triggers on
pydantic-ai's
+``ModelAPIError`` family, which includes ``ModelHTTPError`` -- raised for any
4xx as well as
+5xx. A malformed prompt is the one error every connection in the chain shares:
the same
+request body goes to each of them, so all reject it alike before the task
finally sees the
+failure -- N requests, N timeouts, and N billable calls for a request that was
never going to
+succeed. An expired key does not cost the same way -- it is per-connection, so
the next
+connection in the chain presents its own credentials and, if they are still
valid, answers
+normally; that is the chain doing its job, not a repeated failure. A
misspelled model name is
+shared across the chain only in the narrower case where the name is bare and
every fallback it
+reaches configures no ``model`` of its own: a bare name that itself embeds a
``:`` (a vendor's
+native id) only forwards to a fallback on the *same* platform, and a name that
already pins a
+platform is never forwarded at all -- see *Configure it on the connection*
above for the full
+forwarding rules. Keep chains short, and put deterministic rules for errors
like these in
+:doc:`retry_policies`.
+
+**Airflow's task-level** ``retries`` **multiplies on top of the chain.** A
task with
+``retries=5`` gets up to six attempts -- the initial attempt plus five retries
-- before
+Airflow marks it failed, and each attempt walks the whole chain again if every
connection
+is still down. Against the three-connection chain in the JSON extra under
*Configure it on
+the connection* above (the primary plus two fallbacks), that is up to 18
upstream calls, not
+3, before the task is finally marked failed.
+
+**A bad fallback connection fails the whole chain, including a healthy
primary.** The
+primary and every fallback are resolved eagerly, before any of them is called,
so a
+misspelled fallback ``conn_id`` or a fallback connection missing its ``model``
raises
+immediately -- the task never reaches the primary, even though the primary
itself would
+have answered fine. Run ``test_connection`` on the primary to catch this
before it costs a
+task; see *Verifying a chain* below.
+
+Verifying a chain
+-----------------
+
+Two checks, neither of which requires waiting for a real outage:
+
+*Test the connection.* ``test_connection`` on the primary resolves every
connection in the
+chain, so a fallback with a missing ``model`` or an unknown connection ID is
reported by name
+there rather than discovered mid-incident. Credential fields a provider class
rejects with a
+``TypeError`` are caught by the hook, which retries with the env-var-based
provider
+constructor and logs a warning either way; if the required env var is also
missing, that
+retry raises ``pydantic_ai.exceptions.UserError``, which ``test_connection``
does surface
+since it wraps the whole resolution in a broad exception handler. What it
cannot show is the
+opposite case: the env var *is* set on the worker, the retry quietly succeeds,
and
+``test_connection`` reports success even though the credentials you configured
on the
+connection were silently ignored -- check the logs for that warning rather
than relying on
+``test_connection`` alone. It also does not call the provider, so a
well-formed but revoked
+key still passes -- that is what the drill below is for.
+
+*Drill it.* Point the primary at an endpoint nothing listens on and run the
Dag. The task
+should still succeed, and the run summary in its log names the model that
answered:
+
+.. code-block:: text
+
+ ::group::LLM run complete: model=claude-haiku-4-5-20251001, requests=1, ...
+
+That line is how a failover is noticed at all — it reports the model that
actually served
+the request, not the chain. Repeat the drill whenever the topology changes.
+
+.. exampleinclude::
/../../ai/src/airflow/providers/common/ai/example_dags/example_llm_fallback.py
+ :language: python
+ :dedent: 0
+ :start-after: [START howto_llm_fallback_connection_driven]
+ :end-before: [END howto_llm_fallback_connection_driven]
+
+Scope
+-----
+
+``fallback_conn_ids`` is currently supported only for the pydantic-ai hooks.
Failover here
+is pydantic-ai's ``FallbackModel``, and the other frameworks do not share that
construct:
+LangChain's nearest equivalent is ``Runnable.with_fallbacks()`` on the object
the hook
+returns, and LlamaIndex has none. Extending the same connection-level contract
to them is
+deliberately left out of this change rather than approximated.
diff --git a/providers/common/ai/docs/retry_policies.rst
b/providers/common/ai/docs/retry_policies.rst
index 17c61e378c6..8d1d5f14156 100644
--- a/providers/common/ai/docs/retry_policies.rst
+++ b/providers/common/ai/docs/retry_policies.rst
@@ -91,6 +91,40 @@ If the LLM call fails (provider down, timeout, bad
credentials), the policy
falls back to ``fallback_rules`` if configured, or to the task's standard
retry behaviour.
+This policy decides *between* attempts. Failing over to another vendor *within*
+an attempt is a separate mechanism on the connection — see
+:doc:`provider_fallback`, which also sets out how the two layers compose.
+
+When the connection also carries a fallback chain
+--------------------------------------------------
+
+``LLMRetryPolicy`` builds its classifier hook from ``llm_conn_id`` without
passing
+``fallback_conn_ids``, so if that connection's extra configures a chain (see
+:doc:`provider_fallback`), the policy inherits it silently -- editing the
connection changes
+retry behaviour with no change to the Dag. Two things follow:
+
+* ``timeout`` stops bounding the whole classification call. pydantic-ai
applies a
+ ``ModelSettings`` timeout to each model in the chain, not to the chain as a
whole, so a
+ 30-second ``timeout`` across a three-connection chain is a 90-second worst
case before the
+ policy falls back to ``fallback_rules``.
+* If every connection in the chain fails, the classification call raises
+ ``pydantic_ai.exceptions.FallbackExceptionGroup``. ``evaluate()`` still
degrades to
+ ``fallback_rules`` correctly -- it catches the broad ``Exception``, and an
exception group is
+ one -- so the only cost here is that the classification is wasted.
+
+Separately, and regardless of this policy: if the connection **the task
itself** uses to call
+the LLM (for example ``llm_conn_id`` on ``LLMOperator`` or ``AgentOperator``)
carries a fallback
+chain, the exception the task raises once that chain is exhausted is
+``pydantic_ai.exceptions.FallbackExceptionGroup``, not the last provider's own
exception.
+``RetryRule`` matches with ``isinstance``, so a rule written as
+``RetryRule(exception=ModelHTTPError, ...)`` -- in ``fallback_rules`` here or
in a plain
+``ExceptionRetryPolicy`` -- stops matching. Match
``pydantic_ai.exceptions.FallbackExceptionGroup``
+explicitly as well; its only common ancestor with ``ModelAPIError`` is
``Exception``, too
+broad to write a rule against. The original per-model exceptions are still
available on
+``FallbackExceptionGroup.exceptions``, but ``RetryRule`` only compares the
top-level
+exception type, so a rule set that told 429s apart from 400s collapses into
one rule once
+the chain is in play.
+
What the model can and cannot do
--------------------------------
diff --git a/providers/common/ai/provider.yaml
b/providers/common/ai/provider.yaml
index c23468202df..b44d3299903 100644
--- a/providers/common/ai/provider.yaml
+++ b/providers/common/ai/provider.yaml
@@ -180,6 +180,15 @@ connection-types:
type:
- string
- 'null'
+ fallback_conn_ids:
+ label: Fallback Connections
+ description: "Connection IDs to fail over to, in order, while this
provider is unavailable."
+ schema:
+ type:
+ - array
+ - 'null'
+ items:
+ type: string
- hook-class-name:
airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIAzureHook
hook-name: "Pydantic AI (Azure OpenAI)"
connection-type: pydanticai_azure
@@ -204,6 +213,15 @@ connection-types:
type:
- string
- 'null'
+ fallback_conn_ids:
+ label: Fallback Connections
+ description: "Connection IDs to fail over to, in order, while this
provider is unavailable."
+ schema:
+ type:
+ - array
+ - 'null'
+ items:
+ type: string
api_version:
label: API Version
description: >-
@@ -237,6 +255,15 @@ connection-types:
type:
- string
- 'null'
+ fallback_conn_ids:
+ label: Fallback Connections
+ description: "Connection IDs to fail over to, in order, while this
provider is unavailable."
+ schema:
+ type:
+ - array
+ - 'null'
+ items:
+ type: string
region_name:
label: AWS Region
description: "AWS region (e.g. us-east-1). Falls back to
AWS_DEFAULT_REGION env var."
@@ -325,6 +352,15 @@ connection-types:
type:
- string
- 'null'
+ fallback_conn_ids:
+ label: Fallback Connections
+ description: "Connection IDs to fail over to, in order, while this
provider is unavailable."
+ schema:
+ type:
+ - array
+ - 'null'
+ items:
+ type: string
project:
label: GCP Project
description: "Google Cloud project ID. Falls back to
GOOGLE_CLOUD_PROJECT env var."
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_fallback.py
b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_fallback.py
new file mode 100644
index 00000000000..752e5f5bd5a
--- /dev/null
+++
b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_llm_fallback.py
@@ -0,0 +1,101 @@
+# 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.
+"""
+Example Dag demonstrating provider failover, and how to drill it.
+
+The Dag never names a fallback provider: ``llm_primary_down`` carries the
chain in
+its extra, so the failover topology belongs to whoever administers the
connections.
+
+Prerequisites:
+ - Connection ``llm_primary_down`` with ``conn_type='pydanticai'``,
+ ``host='http://127.0.0.1:9/v1'`` (a port nothing listens on, standing in
for an
+ outage), ``password=<any value>``, and
+ ``extra='{"model": "openai:gpt-4o-mini", "fallback_conn_ids":
["llm_fallback"]}'``
+ - Connection ``llm_primary_down_no_chain``: same as ``llm_primary_down`` but
with
+ ``extra='{"model": "openai:gpt-4o-mini"}'`` (no ``fallback_conn_ids``) --
used by
+ the operator-argument and hook-argument Dags, so the chain visibly comes
from the
+ task instead of the connection
+ - Connection ``llm_fallback`` with ``conn_type='pydanticai'``,
+ ``password=<API key>``, ``extra='{"model":
"anthropic:claude-haiku-4-5-20251001"}'``
+ - ``pip install apache-airflow-providers-common-ai[anthropic]``
+
+Run it as a drill: the task succeeds, and its log reports
``model=claude-haiku-...``
+rather than the primary's model. Point ``llm_primary_down`` at a working
endpoint
+and the same log names the primary instead -- that difference is the evidence
the
+chain is live, and it is the check to repeat whenever the topology changes.
+"""
+
+from __future__ import annotations
+
+from airflow.providers.common.ai.operators.llm import LLMOperator
+from airflow.providers.common.compat.sdk import dag, task
+
+# [START howto_llm_fallback_connection_driven]
+
+
+@dag(catchup=False, tags=["example", "fallback", "llm"])
+def example_llm_fallback():
+ LLMOperator(
+ task_id="summarize_through_the_chain",
+ prompt="Summarize the key findings from the Q4 earnings report.",
+ llm_conn_id="llm_primary_down",
+ system_prompt="You are a financial analyst. Be concise.",
+ )
+
+
+example_llm_fallback()
+
+# [END howto_llm_fallback_connection_driven]
+
+
+# [START howto_llm_fallback_operator_argument]
+@dag(catchup=False, tags=["example", "fallback", "llm"])
+def example_llm_fallback_operator_argument():
+ """The chain lives on the task, not the connection -- it overrides any
chain in extra."""
+ LLMOperator(
+ task_id="summarize_with_an_operator_owned_chain",
+ prompt="Summarize the key findings from the Q4 earnings report.",
+ llm_conn_id="llm_primary_down_no_chain",
+ fallback_conn_ids=["llm_fallback"],
+ system_prompt="You are a financial analyst. Be concise.",
+ )
+
+
+example_llm_fallback_operator_argument()
+# [END howto_llm_fallback_operator_argument]
+
+
+# [START howto_llm_fallback_hook_argument]
+@dag(catchup=False, tags=["example", "fallback", "llm"])
+def example_llm_fallback_explicit_chain():
+ @task
+ def classify_with_an_explicit_chain() -> str:
+ """Build the chain in code, for a task that owns its own failover
order."""
+ from airflow.providers.common.ai.hooks.pydantic_ai import
PydanticAIHook
+
+ hook = PydanticAIHook(
+ llm_conn_id="llm_primary_down_no_chain",
+ fallback_conn_ids=["llm_fallback"],
+ )
+ agent = hook.create_agent(instructions="Reply with a single word.")
+ return agent.run_sync("Is a raven a bird?").output
+
+ classify_with_an_explicit_chain()
+
+
+example_llm_fallback_explicit_chain()
+# [END howto_llm_fallback_hook_argument]
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/get_provider_info.py
b/providers/common/ai/src/airflow/providers/common/ai/get_provider_info.py
index a964ee6e823..a5cf51b4826 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/get_provider_info.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/get_provider_info.py
@@ -149,7 +149,12 @@ def get_provider_info():
"label": "Model",
"description": "Model in provider:name format (e.g.
anthropic:claude-sonnet-5, openai:gpt-5)",
"schema": {"type": ["string", "null"]},
- }
+ },
+ "fallback_conn_ids": {
+ "label": "Fallback Connections",
+ "description": "Connection IDs to fail over to, in
order, while this provider is unavailable.",
+ "schema": {"type": ["array", "null"], "items":
{"type": "string"}},
+ },
},
},
{
@@ -171,6 +176,11 @@ def get_provider_info():
"description": "Azure model identifier (e.g.
azure:gpt-4o)",
"schema": {"type": ["string", "null"]},
},
+ "fallback_conn_ids": {
+ "label": "Fallback Connections",
+ "description": "Connection IDs to fail over to, in
order, while this provider is unavailable.",
+ "schema": {"type": ["array", "null"], "items":
{"type": "string"}},
+ },
"api_version": {
"label": "API Version",
"description": "Azure OpenAI API version (e.g.
2024-07-01-preview). Set when the endpoint path does not end in /v1 and the
host is not *.models.ai.azure.com. Falls back to OPENAI_API_VERSION.",
@@ -196,6 +206,11 @@ def get_provider_info():
"description": "Bedrock model identifier (e.g.
bedrock:us.anthropic.claude-opus-4-5)",
"schema": {"type": ["string", "null"]},
},
+ "fallback_conn_ids": {
+ "label": "Fallback Connections",
+ "description": "Connection IDs to fail over to, in
order, while this provider is unavailable.",
+ "schema": {"type": ["array", "null"], "items":
{"type": "string"}},
+ },
"region_name": {
"label": "AWS Region",
"description": "AWS region (e.g. us-east-1). Falls
back to AWS_DEFAULT_REGION env var.",
@@ -261,6 +276,11 @@ def get_provider_info():
"description": "Google model identifier (e.g.
google-cloud:gemini-2.0-flash)",
"schema": {"type": ["string", "null"]},
},
+ "fallback_conn_ids": {
+ "label": "Fallback Connections",
+ "description": "Connection IDs to fail over to, in
order, while this provider is unavailable.",
+ "schema": {"type": ["array", "null"], "items":
{"type": "string"}},
+ },
"project": {
"label": "GCP Project",
"description": "Google Cloud project ID. Falls back to
GOOGLE_CLOUD_PROJECT env var.",
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 c24672cc913..9d9e15d0c16 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
@@ -16,11 +16,14 @@
# under the License.
from __future__ import annotations
+import re
from pathlib import Path
from typing import TYPE_CHECKING, Any, TypeVar, overload
from pydantic_ai import Agent
+from pydantic_ai.exceptions import ModelAPIError
from pydantic_ai.models import infer_model
+from pydantic_ai.models.fallback import FallbackModel
from pydantic_ai.providers import infer_provider, infer_provider_class
from airflow.providers.common.ai.observability import
genai_instrumentation_settings
@@ -33,12 +36,56 @@ OutputT = TypeVar("OutputT")
# do not auto-enable it either").
_UNSET: Any = object()
+FALLBACK_CONN_IDS_EXTRA_KEY = "fallback_conn_ids"
+
if TYPE_CHECKING:
from pydantic_ai.models import KnownModelName, Model
from airflow.providers.common.compat.sdk import Connection
+def _has_recognized_provider_prefix(model_name: str) -> bool:
+ """
+ Return whether the segment before the first ``:`` in *model_name* is a
pydantic-ai provider.
+
+ A ``:`` alone cannot tell a "provider:model" string apart from a bare
model id that
+ happens to contain a ``:`` of its own -- some vendors' native model ids do
(e.g.
+ Bedrock's version-suffixed ``us.anthropic.claude-opus-4-6-v1:0``). Only a
segment that
+ ``infer_provider_class`` actually recognizes counts as a platform prefix.
+ """
+ prefix, sep, _ = model_name.partition(":")
+ if not sep:
+ return False
+ try:
+ infer_provider_class(prefix)
+ except ImportError:
+ return True # recognized provider; its optional dependency just isn't
installed
+ except ValueError:
+ return False
+ return True
+
+
+_PROVIDER_SLUG_RE = re.compile(r"^[a-z][a-z0-9]*(-[a-z0-9]+)*$")
+
+
+def _looks_like_unrecognized_provider_prefix(prefix: str) -> bool:
+ """
+ Return whether *prefix* has the shape of a plausible-but-wrong provider
name.
+
+ Only called after ``_has_recognized_provider_prefix`` has already said the
segment
+ isn't a real provider. A short, hyphenated, all-lowercase slug
(``"google-vertex"``,
+ ``"google-gla"``, a typo like ``"openi"``) is the shape of something the
user meant as
+ a provider prefix. A vendor's own dotted native id (Bedrock's
+ ``"us.anthropic.claude-opus-4-6-v1"``) never matches -- the ``.`` rules it
out -- so this
+ does not fire for the legitimate embedded-colon case.
+
+ This is a heuristic, not an exhaustive classifier: a prefix with uppercase
letters,
+ underscores, a leading digit, or a stray ``.`` of its own will silently
skip the
+ warning even if it was meant as a typo'd provider name.
+ """
+ return bool(_PROVIDER_SLUG_RE.match(prefix))
+
+
class PydanticAIHook(BaseHook):
"""
Hook for LLM access via pydantic-ai.
@@ -53,22 +100,48 @@ class PydanticAIHook(BaseHook):
Connection fields:
- **password**: API key
- **host**: Base URL (optional, e.g. ``https://api.openai.com/v1``)
- - **extra** JSON: ``{"model": "openai:gpt-5.6-sol"}``
+ - **extra** JSON: ``{"model": "openai:gpt-5.6-sol",
+ "fallback_conn_ids": ["anthropic_prod", "bedrock_dr"]}``
:param llm_conn_id: Airflow connection ID for the LLM provider.
- :param model_id: Model identifier in ``provider:model`` format (e.g.
``"openai:gpt-5.6-sol"``).
- Overrides the model stored in the connection's extra field.
+ :param model_id: Model identifier. A name whose segment before the first
``:``
+ is itself a pydantic-ai provider (e.g. ``"openai:gpt-5.6-sol"``) pins
the
+ platform and is used verbatim -- a plain ``:`` alone is not enough,
since
+ some vendors' native model ids contain one of their own (e.g. Bedrock's
+ version-suffixed ``"us.anthropic.claude-opus-4-6-v1:0"``, which is
still a
+ *bare* name here). A bare name is resolved against this connection's
own
+ platform: vendor subclasses (:class:`PydanticAIAzureHook`,
+ :class:`PydanticAIBedrockHook`, :class:`PydanticAIVertexHook`) each
default
+ to their own platform via :attr:`model_provider`; the generic
connection
+ type has none, so a bare name there raises ``ValueError`` instead of
+ reaching pydantic-ai's own, less actionable ``Unknown model`` error.
+ Overrides the model stored in the connection's extra field. Whichever
of
+ the two configures the primary's model is forwarded (only while still
+ bare) down the fallback chain -- see :meth:`_resolve_fallback_models`.
+ :param fallback_conn_ids: Connection IDs to fail over to, in order, when
the
+ primary provider is unavailable. Overrides the ``fallback_conn_ids``
+ list stored in the connection's extra field; pass an empty list to
+ disable a chain configured there. Blank or whitespace-only entries
+ (including a trailing blank line from the Fallback Connections
textarea)
+ are dropped; a chain left entirely blank is treated the same as passing
+ ``[]``. Each entry may point at any ``pydanticai*`` connection type,
so
+ the chain can span providers (for example OpenAI, then Bedrock). See
+ :meth:`get_conn` for the failover semantics and their cost.
"""
conn_name_attr = "llm_conn_id"
default_conn_name = "pydanticai_default"
conn_type = "pydanticai"
hook_name = "Pydantic AI"
+ # Platform to prefix a bare model_id with (e.g. "azure"); None for the
generic
+ # connection type, which has no platform of its own. Vendor subclasses
override this.
+ model_provider: str | None = None
def __init__(
self,
llm_conn_id: str | None = None,
model_id: str | None = None,
+ fallback_conn_ids: list[str] | None = None,
**kwargs,
) -> None:
super().__init__(**kwargs)
@@ -78,9 +151,12 @@ class PydanticAIHook(BaseHook):
# argument values at class-definition time.
self.llm_conn_id = llm_conn_id if llm_conn_id is not None else
self.default_conn_name
self.model_id = model_id
+ # ``None`` means "not configured here, read the connection's extra";
+ # an empty list means "explicitly no fallbacks", overriding the extra.
+ self.fallback_conn_ids = fallback_conn_ids
self._model: Model | None = None
self._conn: Connection | None = None
- self._conn_extra_dejson: dict[str, Any] | None = None
+ self._conn_extra_dejson: dict[str, Any] = {}
@staticmethod
def get_ui_field_behaviour() -> dict[str, Any]:
@@ -127,11 +203,31 @@ class PydanticAIHook(BaseHook):
kwargs["base_url"] = base_url
return kwargs
+ def _get_conn_and_extra(self) -> tuple[Connection, dict[str, Any]]:
+ """Return this hook's connection and its deserialized extra, fetching
at most once."""
+ if self._conn is None:
+ self._conn = self.get_connection(self.llm_conn_id)
+ self._conn_extra_dejson = self._conn.extra_dejson
+ return self._conn, self._conn_extra_dejson
+
+ def _seed_connection(self, conn: Connection) -> None:
+ """
+ Prime this hook's connection cache with an already-fetched
``Connection``.
+
+ Used by :meth:`_resolve_fallback_models`, which must call
``conn.get_hook()`` to
+ dispatch a fallback's hook class from its ``conn_type`` and so already
holds the
+ ``Connection`` that call built the hook from. Without this,
:meth:`_get_conn_and_extra`
+ would fetch that same connection a second time the first time it runs,
doubling the
+ Execution API round trips a fallback chain costs.
+ """
+ self._conn = conn
+ self._conn_extra_dejson = conn.extra_dejson
+
def get_conn(self) -> Model:
"""
Return a configured pydantic-ai ``Model``.
- Resolution order:
+ Resolution order for this hook's own connection:
1. **Explicit credentials** — when :meth:`_get_provider_kwargs` returns
a non-empty dict the provider class is instantiated with those
kwargs
@@ -139,21 +235,151 @@ class PydanticAIHook(BaseHook):
2. **Default resolution** — delegates to pydantic-ai ``infer_model``
which reads standard env vars (``OPENAI_API_KEY``, ``AWS_PROFILE``,
…).
+ A bare ``model_id`` (one with no recognized platform prefix) is
qualified with
+ this connection's own platform before either of the above -- see the
class
+ docstring's ``model_id`` entry for the resolution and
fallback-forwarding rules.
+
+ When ``fallback_conn_ids`` is configured (on the hook or in the
+ connection's extra) the resolved models are wrapped in a pydantic-ai
+ ``FallbackModel``, so a provider outage moves to the next connection
+ *within the same task attempt* instead of failing the task.
+
+ Two costs of that wrapping are worth knowing before configuring a long
+ chain. A ``timeout`` in ``ModelSettings`` is applied by pydantic-ai to
+ every model in the chain rather than to the chain as a whole, so the
+ worst-case wait is the timeout multiplied by the number of connections.
+ And there is no circuit breaker: every call retries the primary first,
+ so during an outage each task instance pays the primary's timeout
again.
+ Keep the primary's timeout short to bound both.
+
The resolved model is cached for the lifetime of this hook instance.
"""
if self._model is not None:
return self._model
- conn = self.get_connection(self.llm_conn_id) if self._conn is None
else self._conn
- extra: dict[str, Any] = (
- conn.extra_dejson if self._conn_extra_dejson is None else
self._conn_extra_dejson
+ model = self._resolve_own_model()
+ fallback_models = self._resolve_fallback_models()
+ # Pin pydantic-ai's own default explicitly: the retry-layers docs pin
this exact
+ # scope (UnexpectedModelBehavior, UsageLimitExceeded, and
ContentFilterError are
+ # deliberately excluded), and pyproject has no upper bound on
pydantic-ai-slim, so
+ # an upstream default change would otherwise move that documented
behaviour silently.
+ self._model = (
+ FallbackModel(model, *fallback_models,
fallback_on=(ModelAPIError,)) if fallback_models else model
+ )
+ return self._model
+
+ def _qualify_model_name(self, model_name: str, *, forwarded_from_conn_id:
str | None = None) -> str:
+ """
+ Prefix a bare model name with this connection's platform.
+
+ Whether *model_name* is bare or already pins a platform is decided by
+ :func:`_has_recognized_provider_prefix`; see the class docstring's
``model_id``
+ entry for the resolution rules -- including why the generic connection
type
+ raises here instead of reaching pydantic-ai's own, less actionable
+ ``Unknown model`` error.
+
+ :param forwarded_from_conn_id: The primary connection's ID, set only
when
+ *model_name* was forwarded down a fallback chain rather than
configured
+ directly on this connection -- see :meth:`_resolve_own_model`.
Used to
+ attribute an unresolvable name to where it actually came from
instead of
+ blaming this (fallback) connection for a name it never set.
+ """
+ if model_name == "test":
+ return model_name
+ if _has_recognized_provider_prefix(model_name):
+ return model_name
+ if self.model_provider is not None:
+ prefix, sep, _ = model_name.partition(":")
+ if sep and _looks_like_unrecognized_provider_prefix(prefix):
+ self.log.warning(
+ "Model name '%s' on connection '%s' contains ':' but its
prefix '%s' is not a "
+ "provider pydantic-ai recognizes; treating the whole
string as a bare %s model id "
+ "and resolving it as '%s:%s'. If '%s' was meant to be a
provider prefix, this looks "
+ "like it might be a typo.",
+ model_name,
+ self.llm_conn_id,
+ prefix,
+ self.model_provider,
+ self.model_provider,
+ model_name,
+ prefix,
+ )
+ return f"{self.model_provider}:{model_name}"
+
+ if forwarded_from_conn_id is not None:
+ raise ValueError(
+ f"Connection '{self.llm_conn_id}' has no default model
provider, so the bare model "
+ f"name '{model_name}' -- forwarded from primary connection
'{forwarded_from_conn_id}' "
+ f"-- cannot be resolved here. Give '{forwarded_from_conn_id}'
a 'provider:model' "
+ f"string, or set an explicit 'model' on '{self.llm_conn_id}'."
+ )
+
+ prefix, sep, _ = model_name.partition(":")
+ if sep:
+ raise ValueError(
+ f"Connection '{self.llm_conn_id}' has no default model
provider, and '{prefix}' is "
+ f"not a provider pydantic-ai recognizes, so '{model_name}'
cannot be resolved. If "
+ "this is a vendor's own model id containing a ':' (e.g. a
Bedrock-style "
+ "version-suffixed id), use a vendor connection type
(Azure/Bedrock/Vertex) instead; "
+ f"if '{prefix}' is meant to be a provider prefix, check it for
a typo."
+ )
+ raise ValueError(
+ f"Connection '{self.llm_conn_id}' has no default model provider,
so the bare model name "
+ f"'{model_name}' cannot be resolved. Use a vendor connection type
(Azure/Bedrock/Vertex) "
+ "or set an explicit 'provider:model' string."
)
- model_name: str | KnownModelName = self.model_id or extra.get("model",
"")
+ def _get_configured_model_name(self) -> str | KnownModelName | None:
+ """Return the model name this connection configures, hook argument
winning over the extra."""
+ if self.model_id:
+ return self.model_id
+ _, extra = self._get_conn_and_extra()
+ return extra.get("model")
+
+ def _resolve_own_model(
+ self,
+ *,
+ forwarded_model_id: str | None = None,
+ forwarded_from_conn_id: str | None = None,
+ forwarded_model_provider: str | None = None,
+ ) -> Model:
+ """
+ Resolve the ``Model`` for this hook's own connection, ignoring any
fallback chain.
+
+ :param forwarded_model_id: The primary connection's configured model
name,
+ forwarded down a fallback chain by
:meth:`_resolve_fallback_models` --
+ see that method's docstring for when a name is eligible to
forward. A
+ name with an embedded ``:`` of its own (a vendor's own native id,
e.g.
+ Bedrock's version-suffixed ``us.anthropic.claude-opus-4-6-v1:0``)
is only
+ forwarded when *forwarded_model_provider* matches this
connection's own
+ :attr:`model_provider` -- that spelling is only meaningful on the
+ platform that produced it.
+ :param forwarded_from_conn_id: The primary connection's ID, for error
messages
+ attributing an unresolvable forwarded name to where it actually
came from.
+ :param forwarded_model_provider: The primary connection's
:attr:`model_provider`;
+ see *forwarded_model_id* above for how it gates forwarding.
+ """
+ conn, extra = self._get_conn_and_extra()
+
+ model_name: str | KnownModelName | None =
self._get_configured_model_name()
+ forwarded = False
+ if (
+ not model_name
+ and forwarded_model_id
+ and not _has_recognized_provider_prefix(forwarded_model_id)
+ and (":" not in forwarded_model_id or forwarded_model_provider ==
self.model_provider)
+ ):
+ model_name = forwarded_model_id
+ forwarded = True
if not model_name:
raise ValueError(
- "No model specified. Set model_id on the hook or the Model
field on the connection."
+ f"No model specified for connection '{self.llm_conn_id}'. Set
model_id on the "
+ "hook or the Model field on the connection."
)
+ model_name = self._qualify_model_name(
+ model_name,
+ forwarded_from_conn_id=forwarded_from_conn_id if forwarded else
None,
+ )
api_key: str | None = conn.password or None
base_url: str | None = conn.host or None
@@ -178,24 +404,119 @@ class PydanticAIHook(BaseHook):
)
return infer_provider(pname)
- self._model = infer_model(model_name,
provider_factory=_provider_factory)
- return self._model
+ return infer_model(model_name, provider_factory=_provider_factory)
- self._model = infer_model(model_name)
- return self._model
+ return infer_model(model_name)
- def _get_conn_if_model_configured(self) -> Model | None:
- """Return the hook model only when the hook or connection explicitly
configures one."""
- if self.model_id:
- return self.get_conn()
+ def _get_fallback_conn_ids(self) -> list[str]:
+ """
+ Return the configured fallback connection IDs, hook argument winning
over the extra.
- conn = self.get_connection(self.llm_conn_id)
- self._conn = conn
- self._conn_extra_dejson = conn.extra_dejson
+ Blank entries (including whitespace-only ones) are dropped and
surviving entries are
+ stripped: the Fallback Connections field renders as a textarea that
splits on newline,
+ and its blur handler only guards against an all-blank value, so a
trailing blank line
+ is what most saved chains actually look like.
+ """
+ if self.fallback_conn_ids is not None:
+ raw: Any = self.fallback_conn_ids
+ else:
+ _, extra = self._get_conn_and_extra()
+ raw = extra.get(FALLBACK_CONN_IDS_EXTRA_KEY)
+ if raw is None:
+ raw = []
+
+ if not isinstance(raw, (list, tuple)) or not all(isinstance(item, str)
for item in raw):
+ raise ValueError(
+ f"{FALLBACK_CONN_IDS_EXTRA_KEY} for connection
'{self.llm_conn_id}' must be a list "
+ f"of connection IDs, got {raw!r}."
+ )
+ return [stripped for item in raw if (stripped := item.strip())]
+
+ def _resolve_fallback_models(self) -> list[Model]:
+ """
+ Resolve one ``Model`` per fallback connection, in the configured order.
+
+ Each connection is resolved through the hook registered for its own
+ ``conn_type``, so a chain can mix providers whose credentials live in
+ different connection fields. The primary's configured model name --
its
+ ``model_id`` argument, or the ``model`` in its own ``extra`` -- is
forwarded
+ to each fallback as a logical model name: a fallback connection with
its own
+ ``model`` in ``extra`` uses that instead, but a fallback with none
falls
+ back to the forwarded name, qualified with *its own* platform prefix.
+ Only a *bare* forwarded name is usable this way -- a forwarded name
that
+ already pins a platform (e.g. ``"openai:gpt-5"``) names a model of the
+ primary's provider, not this fallback's, so it is not applied; that
+ fallback still raises "no model specified" unless its own ``extra``
sets
+ a ``model``. Whether a name already pins a platform is decided by
+ :func:`_has_recognized_provider_prefix`, not by whether it merely
contains a ``:``.
+ """
+ fallback_conn_ids = self._get_fallback_conn_ids()
+ if not fallback_conn_ids:
+ return []
+
+ forwarded_model_id = self._get_configured_model_name()
+
+ self.log.info("Resolving LLM fallback chain: %s", " ->
".join([self.llm_conn_id, *fallback_conn_ids]))
+
+ models: list[Model] = []
+ seen: set[str] = set()
+ for conn_id in fallback_conn_ids:
+ if conn_id == self.llm_conn_id:
+ raise ValueError(
+ f"Fallback chain for connection '{self.llm_conn_id}' lists
the primary "
+ "connection as one of its own fallbacks; every fallback
must differ from "
+ "the primary."
+ )
+ if conn_id in seen:
+ raise ValueError(
+ f"Fallback chain for connection '{self.llm_conn_id}' lists
'{conn_id}' more "
+ "than once; every entry must be distinct."
+ )
+ seen.add(conn_id)
+
+ # ``PydanticAIHook.get_hook(conn_id)`` would fetch this connection
twice: once
+ # inside itself and once more the first time the new hook's own
+ # ``_get_conn_and_extra`` runs (see ``_seed_connection``). Fetch
it once here and
+ # dispatch the hook class from it directly instead -- this is
exactly what
+ # ``BaseHook.get_hook`` does internally, so the result still isn't
constrained to
+ # this class and the type has to be checked here.
+ conn = PydanticAIHook.get_connection(conn_id)
+ hook = conn.get_hook()
+ if not isinstance(hook, PydanticAIHook):
+ raise ValueError(
+ f"Fallback connection '{conn_id}' resolves to
{type(hook).__name__}, which is "
+ "not a PydanticAIHook. Only pydanticai connection types
can be used as "
+ f"fallbacks for '{self.llm_conn_id}'."
+ )
+ hook._seed_connection(conn)
+ if hook._get_fallback_conn_ids():
+ raise ValueError(
+ f"Fallback connection '{conn_id}' declares its own "
+ f"{FALLBACK_CONN_IDS_EXTRA_KEY}. Chains are not resolved
recursively -- list "
+ f"every provider directly on '{self.llm_conn_id}' instead."
+ )
+ models.append(
+ hook._resolve_own_model(
+ forwarded_model_id=forwarded_model_id,
+ forwarded_from_conn_id=self.llm_conn_id,
+ forwarded_model_provider=self.model_provider,
+ )
+ )
+
+ return models
- if self._conn_extra_dejson.get("model"):
+ def _get_conn_if_model_configured(self) -> Model | None:
+ """Return the hook model only when the hook or connection explicitly
configures one."""
+ if self._get_configured_model_name():
return self.get_conn()
+ if self._get_fallback_conn_ids():
+ raise ValueError(
+ f"A fallback chain is configured for '{self.llm_conn_id}' but
no model is set. "
+ "A fallback chain needs an explicit primary model -- set the
Model field on the "
+ "connection or model_id on the hook. (A model taken from an
agent spec file "
+ "cannot be wrapped in a fallback chain.)"
+ )
return None
@overload
@@ -248,8 +569,11 @@ class PydanticAIHook(BaseHook):
use only the file value.
:param spec_file: Path to a YAML or JSON ``AgentSpec`` file. When
supplied,
delegates to ``Agent.from_file``. If ``model_id`` or the
connection's
- ``model`` extra is set, that model is passed to pydantic-ai;
otherwise
- the spec file's ``model`` is used.
+ ``model`` extra is set, that model is passed to pydantic-ai;
otherwise the
+ spec file's ``model`` is used. A connection that declares
+ ``fallback_conn_ids`` but no ``model`` raises ``ValueError``
instead: a model
+ resolved from the spec file cannot be wrapped in a fallback chain,
so the
+ chain would otherwise be dropped silently.
:param agent_kwargs: Additional keyword arguments passed to the Agent
constructor.
"""
# ``instrument`` is no longer an ``Agent()`` / ``Agent.from_file()``
@@ -290,10 +614,16 @@ class PydanticAIHook(BaseHook):
"""
Test connection by resolving the model.
- Validates that the model string is valid and the provider class can be
- instantiated with the supplied credentials. Does NOT make an LLM API
- call — that would be expensive and fail for reasons unrelated to
- connectivity (quotas, billing, rate limits).
+ A success here can come from this connection's own credentials, or --
when a
+ provider class rejects them with a ``TypeError`` -- from a silent
retry against
+ the standard environment variables, which ignores those credentials
entirely.
+ See :doc:`/provider_fallback`'s *Verifying a chain* section for how to
tell the
+ two apart. Does NOT make an LLM API call — that would be expensive and
fail for
+ reasons unrelated to connectivity (quotas, billing, rate limits).
+
+ Every connection in ``fallback_conn_ids`` is resolved too, so a
+ misconfigured fallback is reported here rather than discovered during
+ the outage it was meant to cover.
"""
try:
self.get_conn()
@@ -324,6 +654,7 @@ class PydanticAIAzureHook(PydanticAIHook):
conn_type = "pydanticai_azure"
default_conn_name = "pydanticai_azure_default"
hook_name = "Pydantic AI (Azure OpenAI)"
+ model_provider = "azure"
@staticmethod
def get_ui_field_behaviour() -> dict[str, Any]:
@@ -392,6 +723,7 @@ class PydanticAIBedrockHook(PydanticAIHook):
conn_type = "pydanticai_bedrock"
default_conn_name = "pydanticai_bedrock_default"
hook_name = "Pydantic AI (AWS Bedrock)"
+ model_provider = "bedrock"
@staticmethod
def get_ui_field_behaviour() -> dict[str, Any]:
@@ -477,13 +809,26 @@ class PydanticAIVertexHook(PydanticAIHook):
model prefix (``google-cloud:`` vs. ``google:``) rather than a
constructor flag, so there is nothing left for this field to control.
+ A bare ``model_id`` (or Extra ``model``) always defaults to Vertex AI
+ (``google-cloud:``) -- this default is **not** inferred from which
+ credential fields are set. ``api_key`` in ``extra`` can mean either the
+ Generative Language API or Vertex API-key auth (see credential order
+ above), so its presence alone cannot tell the two platforms apart; guessing
+ would risk silently authenticating against the wrong endpoint. To use the
+ Generative Language API, set an explicit ``google:``-prefixed model id (on
+ the hook or the connection's ``model`` extra) -- that spelling already
+ works today.
+
:param llm_conn_id: Airflow connection ID.
- :param model_id: Model identifier, e.g.
``"google-cloud:gemini-2.0-flash"``.
+ :param model_id: Model identifier, e.g.
``"google-cloud:gemini-2.0-flash"``. A
+ bare name (e.g. ``"gemini-2.0-flash"``) defaults to Vertex AI; prefix
with
+ ``google:`` for the Generative Language API.
"""
conn_type = "pydanticai_vertex"
default_conn_name = "pydanticai_vertex_default"
hook_name = "Pydantic AI (Google Vertex AI)"
+ model_provider = "google-cloud"
@staticmethod
def get_ui_field_behaviour() -> dict[str, Any]:
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 e5d1212dad5..6d3b1f71fa3 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
@@ -141,6 +141,13 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
:param llm_conn_id: Connection ID for the LLM provider.
:param model_id: Model identifier (e.g. ``"openai:gpt-5"``).
Overrides the model stored in the connection's extra field.
+ :param fallback_conn_ids: Connection IDs to fail over to, in order, when
+ the primary provider is unavailable. Overrides the
``fallback_conn_ids``
+ set in the connection's extra field. ``None`` (default) reads the
+ connection's own extra field; an explicit ``[]`` disables a chain
+ configured there. See
+ :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`
+ for how blank entries in the list are dropped.
:param system_prompt: System-level instructions for the agent.
:param output_type: Expected output type. Default ``str``. Set to a
Pydantic
``BaseModel`` subclass for structured output; the model instance is
@@ -255,6 +262,7 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
"prompt",
"llm_conn_id",
"model_id",
+ "fallback_conn_ids",
"system_prompt",
"agent_params",
"message_history",
@@ -269,6 +277,7 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
prompt: str,
llm_conn_id: str,
model_id: str | None = None,
+ fallback_conn_ids: list[str] | None = None,
system_prompt: str = "",
output_type: type = str,
toolsets: list[AbstractToolset] | None = None,
@@ -291,6 +300,7 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
self.prompt = prompt
self.llm_conn_id = llm_conn_id
self.model_id = model_id
+ self.fallback_conn_ids = fallback_conn_ids
self.system_prompt = system_prompt
self.output_type = output_type
self.serialize_output = serialize_output
@@ -350,6 +360,7 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
"""Return PydanticAIHook for the configured LLM connection."""
hook_params = {
"model_id": self.model_id,
+ "fallback_conn_ids": self.fallback_conn_ids,
}
return PydanticAIHook.get_hook(self.llm_conn_id,
hook_params=hook_params)
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
index ede18548abc..17e9da95b90 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
@@ -69,6 +69,13 @@ class LLMOperator(BaseOperator, LLMApprovalMixin):
:param llm_conn_id: Connection ID for the LLM provider.
:param model_id: Model identifier (e.g. ``"openai:gpt-5"``).
Overrides the model stored in the connection's extra field.
+ :param fallback_conn_ids: Connection IDs to fail over to, in order, when
+ the primary provider is unavailable. Overrides the
``fallback_conn_ids``
+ set in the connection's extra field. ``None`` (default) reads the
+ connection's own extra field; an explicit ``[]`` disables a chain
+ configured there. See
+ :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`
+ for how blank entries in the list are dropped.
:param system_prompt: System-level instructions for the LLM agent.
:param output_type: Expected output type. Default ``str``. Set to a
Pydantic
``BaseModel`` subclass for structured output; the model instance is
@@ -137,6 +144,7 @@ class LLMOperator(BaseOperator, LLMApprovalMixin):
"prompt",
"llm_conn_id",
"model_id",
+ "fallback_conn_ids",
"system_prompt",
"agent_params",
"usage_limits",
@@ -148,6 +156,7 @@ class LLMOperator(BaseOperator, LLMApprovalMixin):
prompt: str,
llm_conn_id: str,
model_id: str | None = None,
+ fallback_conn_ids: list[str] | None = None,
system_prompt: str = "",
output_type: type = str,
agent_params: dict[str, Any] | None = None,
@@ -165,6 +174,7 @@ class LLMOperator(BaseOperator, LLMApprovalMixin):
self.prompt = prompt
self.llm_conn_id = llm_conn_id
self.model_id = model_id
+ self.fallback_conn_ids = fallback_conn_ids
self.system_prompt = system_prompt
self.output_type = output_type
self.serialize_output = serialize_output
@@ -246,6 +256,7 @@ class LLMOperator(BaseOperator, LLMApprovalMixin):
"""
hook_params = {
"model_id": self.model_id,
+ "fallback_conn_ids": self.fallback_conn_ids,
}
return PydanticAIHook.get_hook(self.llm_conn_id,
hook_params=hook_params)
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py
index c7b68f9d22a..3541efa85f9 100644
---
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py
+++
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_branch.py
@@ -45,6 +45,13 @@ class LLMBranchOperator(LLMOperator, BranchMixIn):
:param llm_conn_id: Connection ID for the LLM provider.
:param model_id: Model identifier (e.g. ``"openai:gpt-5"``).
Overrides the model stored in the connection's extra field.
+ :param fallback_conn_ids: Connection IDs to fail over to, in order, when
+ the primary provider is unavailable. Overrides the
``fallback_conn_ids``
+ set in the connection's extra field. ``None`` (default) reads the
+ connection's own extra field; an explicit ``[]`` disables a chain
+ configured there. See
+ :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`
+ for how blank entries in the list are dropped.
:param system_prompt: System-level instructions for the LLM agent.
:param allow_multiple_branches: When ``False`` (default) the LLM returns a
single task ID. When ``True`` the LLM may return one or more task IDs.
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_file_analysis.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_file_analysis.py
index 31d4635e654..b6a7593886b 100644
---
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_file_analysis.py
+++
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_file_analysis.py
@@ -47,6 +47,13 @@ class LLMFileAnalysisOperator(LLMOperator):
:param llm_conn_id: Connection ID for the LLM provider.
:param model_id: Model identifier (e.g. ``"openai:gpt-5"``).
Overrides the model stored in the connection's extra field.
+ :param fallback_conn_ids: Connection IDs to fail over to, in order, when
+ the primary provider is unavailable. Overrides the
``fallback_conn_ids``
+ set in the connection's extra field. ``None`` (default) reads the
+ connection's own extra field; an explicit ``[]`` disables a chain
+ configured there. See
+ :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`
+ for how blank entries in the list are dropped.
:param system_prompt: Additional instructions appended to the built-in
file-analysis system prompt.
:param agent_params: Additional keyword arguments passed to the pydantic-ai
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_schema_compare.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_schema_compare.py
index ccb4a62cb8a..82cacb2cac0 100644
---
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_schema_compare.py
+++
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_schema_compare.py
@@ -97,6 +97,13 @@ class LLMSchemaCompareOperator(LLMOperator):
:param prompt: Instructions for the LLM on what to compare and flag.
:param llm_conn_id: Connection ID for the LLM provider.
:param model_id: Model identifier (e.g. ``"openai:gpt-5"``).
+ :param fallback_conn_ids: Connection IDs to fail over to, in order, when
+ the primary provider is unavailable. Overrides the
``fallback_conn_ids``
+ set in the connection's extra field. ``None`` (default) reads the
+ connection's own extra field; an explicit ``[]`` disables a chain
+ configured there. See
+ :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`
+ for how blank entries in the list are dropped.
:param system_prompt: Instructions included in the LLM system prompt.
Defaults to
``DEFAULT_SYSTEM_PROMPT`` which contains cross-system type
equivalences and
severity definitions. Passing a value **replaces** the default system
prompt
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_sql.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_sql.py
index 96c6453e092..2caba10b9ef 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/operators/llm_sql.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llm_sql.py
@@ -65,6 +65,13 @@ class LLMSQLQueryOperator(LLMOperator):
:param llm_conn_id: Connection ID for the LLM provider.
:param model_id: Model identifier (e.g. ``"openai:gpt-4o"``).
Overrides the model stored in the connection's extra field.
+ :param fallback_conn_ids: Connection IDs to fail over to, in order, when
+ the primary provider is unavailable. Overrides the
``fallback_conn_ids``
+ set in the connection's extra field. ``None`` (default) reads the
+ connection's own extra field; an explicit ``[]`` disables a chain
+ configured there. See
+ :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`
+ for how blank entries in the list are dropped.
:param system_prompt: Additional instructions appended to the built-in SQL
safety prompt. Use for domain-specific guidance.
:param agent_params: Additional keyword arguments passed to the pydantic-ai
diff --git a/providers/common/ai/tests/unit/common/ai/decorators/test_llm.py
b/providers/common/ai/tests/unit/common/ai/decorators/test_llm.py
index e68d6ce6a78..a66a4006bc3 100644
--- a/providers/common/ai/tests/unit/common/ai/decorators/test_llm.py
+++ b/providers/common/ai/tests/unit/common/ai/decorators/test_llm.py
@@ -47,6 +47,25 @@ class TestLLMDecoratedOperator:
assert op.prompt == "Summarize this text"
mock_agent.run_sync.assert_called_once_with("Summarize this text",
usage_limits=None)
+ @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook",
autospec=True)
+ def test_execute_forwards_fallback_conn_ids_to_hook(self, mock_hook_cls,
make_mock_run_result):
+ """``fallback_conn_ids`` is accepted as a decorator kwarg without a
per-decorator code change."""
+ mock_agent = MagicMock(spec=["run_sync"])
+ mock_agent.run_sync.return_value = make_mock_run_result("ok")
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
mock_agent
+
+ op = _LLMDecoratedOperator(
+ task_id="test",
+ python_callable=lambda: "p",
+ llm_conn_id="my_llm",
+ fallback_conn_ids=["conn_a", "conn_b"],
+ )
+ op.execute(context={})
+
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids":
["conn_a", "conn_b"]}
+ )
+
@pytest.mark.parametrize(
"return_value",
[42, "", " ", None, b"bytes", bytearray(b"x"), [], ()],
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 d4ef9a46de6..9db95fce059 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
@@ -18,13 +18,18 @@ from __future__ import annotations
import contextlib
import json
+import logging
import re
import sys
from pathlib import Path
-from unittest.mock import MagicMock, patch
+from unittest.mock import MagicMock, call, patch
import pytest
+from pydantic_ai import Agent
+from pydantic_ai.exceptions import FallbackExceptionGroup, ModelAPIError
from pydantic_ai.models import Model
+from pydantic_ai.models.fallback import FallbackModel
+from pydantic_ai.models.function import FunctionModel
from pydantic_ai.models.test import TestModel
from pydantic_ai.providers import infer_provider_class
@@ -35,7 +40,10 @@ from airflow.providers.common.ai.hooks.pydantic_ai import (
PydanticAIBedrockHook,
PydanticAIHook,
PydanticAIVertexHook,
+ _has_recognized_provider_prefix,
+ _looks_like_unrecognized_provider_prefix,
)
+from airflow.providers.common.compat.sdk import AirflowNotFoundException
# Matches the `google...` provider key pydantic-ai expects before the
`:model-name`
# separator, e.g. "google-cloud" out of "google-cloud:gemini-2.0-flash".
@@ -240,6 +248,864 @@ class TestPydanticAIHookGetConn:
mock_infer_model.assert_called_once()
+class _ConnRegistry:
+ """
+ In-memory stand-in for connection and hook lookup.
+
+ ``_resolve_fallback_models`` calls ``PydanticAIHook.get_connection`` and
then
+ ``Connection.get_hook`` on the result, both of which need the metadata DB
and
+ provider discovery; this resolves both from a dict instead.
+ """
+
+ def __init__(self) -> None:
+ self.conns: dict[str, Connection] = {}
+ self.hook_classes: dict[str, type[PydanticAIHook]] = {}
+
+ def add(
+ self,
+ conn_id: str,
+ *,
+ conn_type: str = "pydanticai",
+ hook_class: type[PydanticAIHook] = PydanticAIHook,
+ password: str | None = None,
+ extra: dict | None = None,
+ ) -> None:
+ self.conns[conn_id] = Connection(
+ conn_id=conn_id,
+ conn_type=conn_type,
+ password=password,
+ extra=json.dumps(extra) if extra else None,
+ )
+ self.hook_classes[conn_id] = hook_class
+
+ def get_connection(self, conn_id: str) -> Connection:
+ try:
+ return self.conns[conn_id]
+ except KeyError:
+ raise AirflowNotFoundException(f"The conn_id `{conn_id}` isn't
defined") from None
+
+ def get_hook(self, conn: Connection, *, hook_params: dict | None = None):
+ """Side effect for the patched ``Connection.get_hook`` -- ``conn`` is
bound as ``self``
+ (via ``autospec=True`` on the patch), so the connection to dispatch
from is this
+ argument's ``conn_id``, not a value the caller passes in.
+ """
+ if conn.conn_id not in self.conns:
+ raise AirflowNotFoundException(f"The conn_id `{conn.conn_id}`
isn't defined")
+ hook_class = self.hook_classes[conn.conn_id]
+ return hook_class(llm_conn_id=conn.conn_id, **(hook_params or {}))
+
+
[email protected]
+def registry():
+ """Patch connection and hook lookup onto a registry the test populates."""
+ reg = _ConnRegistry()
+ with (
+ patch.object(PydanticAIHook, "get_connection",
side_effect=reg.get_connection),
+ # autospec=True is required here: it's what makes the mock bind `self`
(the
+ # Connection instance `.get_hook()` was called on) as the side
effect's first
+ # argument -- without it, `conn.get_hook()` calls the mock with zero
arguments and
+ # `reg.get_hook` would have no way to know which connection dispatched
it.
+ patch.object(Connection, "get_hook", side_effect=reg.get_hook,
autospec=True),
+ ):
+ yield reg
+
+
+class _InferModelStub:
+ """Resolve every model string to its own recognisable model, and record
how it was built."""
+
+ def __init__(self, mock: MagicMock) -> None:
+ self.mock = mock
+ self.models: dict[str, MagicMock] = {}
+
+ def __call__(self, model_name: str, **kwargs) -> MagicMock:
+ return self.models.setdefault(model_name, MagicMock(spec=Model,
name=model_name))
+
+ def provider_kwargs_for(self, model_name: str, infer_provider_class:
MagicMock) -> dict:
+ """Return the kwargs the provider for *model_name* would be
constructed with."""
+ factory = next(
+ call.kwargs["provider_factory"] for call in
self.mock.call_args_list if call.args[0] == model_name
+ )
+ infer_provider_class.return_value.reset_mock()
+ factory(model_name.split(":")[0])
+ return infer_provider_class.return_value.call_args.kwargs
+
+
[email protected]
+def infer_model_stub():
+ """Patch ``infer_model`` so tests can tell the models of a chain apart."""
+ with patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_model",
autospec=True) as mock:
+ stub = _InferModelStub(mock)
+ mock.side_effect = stub
+ yield stub
+
+
+class TestPydanticAIHookModelProviderResolution:
+ """Bare model names get qualified with a connection's own platform
prefix."""
+
+ @pytest.mark.parametrize(
+ ("hook_class", "conn_type", "prefix"),
+ [
+ pytest.param(PydanticAIAzureHook, "pydanticai_azure", "azure",
id="azure"),
+ pytest.param(PydanticAIBedrockHook, "pydanticai_bedrock",
"bedrock", id="bedrock"),
+ pytest.param(PydanticAIVertexHook, "pydanticai_vertex",
"google-cloud", id="vertex"),
+ ],
+ )
+ def test_bare_model_id_gets_platform_prefix(
+ self, registry, infer_model_stub, hook_class, conn_type, prefix
+ ):
+ registry.add("primary", conn_type=conn_type, hook_class=hook_class,
extra={"model": "foo"})
+ hook = hook_class(llm_conn_id="primary")
+
+ assert hook.get_conn() is infer_model_stub.models[f"{prefix}:foo"]
+
+ @pytest.mark.parametrize(
+ ("hook_class", "conn_type"),
+ [
+ pytest.param(PydanticAIHook, "pydanticai", id="generic"),
+ pytest.param(PydanticAIAzureHook, "pydanticai_azure", id="azure"),
+ pytest.param(PydanticAIBedrockHook, "pydanticai_bedrock",
id="bedrock"),
+ pytest.param(PydanticAIVertexHook, "pydanticai_vertex",
id="vertex"),
+ ],
+ )
+ def test_test_sentinel_is_never_platform_prefixed(
+ self, registry, infer_model_stub, hook_class, conn_type
+ ):
+ """The literal ``"test"`` model name is pydantic-ai's dry-run sentinel
and must reach
+ ``infer_model`` unprefixed on every connection type, including ones
with a platform.
+
+ Mutation canary: dropping the ``model_name == "test"`` early return in
+ ``_qualify_model_name`` makes this raise on the generic hook (no
default model
+ provider) and resolve to e.g. ``"azure:test"`` on the vendor hooks --
either way the
+ ``is infer_model_stub.models["test"]`` identity assertion fails.
+ """
+ registry.add("primary", conn_type=conn_type, hook_class=hook_class,
extra={"model": "test"})
+ hook = hook_class(llm_conn_id="primary")
+
+ assert hook.get_conn() is infer_model_stub.models["test"]
+
+ def test_prefixed_model_id_used_verbatim(self, registry, infer_model_stub):
+ """A name that already contains ``:`` pins its own platform and is
never re-prefixed."""
+ registry.add(
+ "primary",
+ conn_type="pydanticai_azure",
+ hook_class=PydanticAIAzureHook,
+ extra={"model": "openai:gpt-4"},
+ )
+ hook = PydanticAIAzureHook(llm_conn_id="primary")
+
+ assert hook.get_conn() is infer_model_stub.models["openai:gpt-4"]
+
+ def test_generic_connection_bare_name_raises_actionable_error(self,
registry, infer_model_stub):
+ """The generic ``pydanticai`` connection type has no platform of its
own."""
+ registry.add("primary", extra={"model": "gpt-4"})
+ hook = PydanticAIHook(llm_conn_id="primary")
+
+ with pytest.raises(ValueError, match="primary") as exc_info:
+ hook.get_conn()
+
+ assert "gpt-4" in str(exc_info.value)
+
+ def test_vertex_bare_model_id_ignores_credential_shape(self, registry,
infer_model_stub):
+ """Vertex's default platform never depends on which credential fields
are set.
+
+ ``api_key`` in this hook's extra can mean either the Generative
Language API or
+ Vertex API-key auth, so it cannot decide the platform -- there is
deliberately no
+ inference here, only the class-level default.
+ """
+ registry.add(
+ "primary",
+ conn_type="pydanticai_vertex",
+ hook_class=PydanticAIVertexHook,
+ extra={"model": "gemini-2.0-flash", "api_key": "some-key"},
+ )
+ hook = PydanticAIVertexHook(llm_conn_id="primary")
+
+ assert hook.get_conn() is
infer_model_stub.models["google-cloud:gemini-2.0-flash"]
+
+ def
test_bedrock_bare_model_id_with_embedded_colon_gets_platform_prefix(self,
registry, infer_model_stub):
+ """A ``:`` alone doesn't pin a platform -- Bedrock's own ids contain
one.
+
+ Bedrock's version-suffixed ids (e.g.
``us.anthropic.claude-opus-4-6-v1:0``) contain
+ a ``:`` that is not a pydantic-ai provider name, so a bare copy of one
must still get
+ the ``bedrock:`` prefix, not be treated as already-qualified.
+
+ Mutation canary: reverting
``_qualify_model_name``/``_has_recognized_provider_prefix``
+ to the old ``":" in model_name`` check makes this resolve to the
unprefixed
+ ``"us.anthropic.claude-opus-4-6-v1:0"`` instead, failing the ``is``
identity assertion
+ (a different key in ``infer_model_stub.models``).
+ """
+ registry.add(
+ "primary",
+ conn_type="pydanticai_bedrock",
+ hook_class=PydanticAIBedrockHook,
+ extra={"model": "us.anthropic.claude-opus-4-6-v1:0"},
+ )
+ hook = PydanticAIBedrockHook(llm_conn_id="primary")
+
+ assert hook.get_conn() is
infer_model_stub.models["bedrock:us.anthropic.claude-opus-4-6-v1:0"]
+
+ def test_prefixed_model_id_with_embedded_colon_used_verbatim(self,
registry, infer_model_stub):
+ """A name already pinning a recognized platform is never re-prefixed,
even with an
+ embedded ``:`` of its own.
+
+ Mutation canary: dropping the ``infer_provider_class`` recognition
check (treating
+ every ``:`` split the same) has no effect on *this* test by itself
since the string
+ already starts with a recognized prefix -- what would catch a
regression here is a
+ mutation that re-adds prefixing unconditionally (e.g. always prepending
+ ``model_provider`` regardless of ``_has_recognized_provider_prefix``'s
result), which
+ would turn the resolved key into
+ ``"bedrock:bedrock:us.anthropic.claude-opus-4-6-v1:0"`` and fail the
identity assertion.
+ """
+ registry.add(
+ "primary",
+ conn_type="pydanticai_bedrock",
+ hook_class=PydanticAIBedrockHook,
+ extra={"model": "bedrock:us.anthropic.claude-opus-4-6-v1:0"},
+ )
+ hook = PydanticAIBedrockHook(llm_conn_id="primary")
+
+ assert hook.get_conn() is
infer_model_stub.models["bedrock:us.anthropic.claude-opus-4-6-v1:0"]
+
+ def
test_generic_connection_unrecognized_prefix_raises_actionable_error(self,
registry, infer_model_stub):
+ """A ``:`` whose left segment isn't a real provider must not slip past
as "prefixed".
+
+ On the generic connection type (no platform of its own) a name like
+ ``"us.anthropic.claude-opus-4-6-v1:0"`` must raise this hook's own
actionable error
+ naming the connection, not be forwarded to pydantic-ai's
``infer_model`` where it
+ would instead raise the less actionable ``UserError: Unknown model``.
+
+ Mutation canary: reverting to the old ``":" in model_name`` check
makes this string
+ look "already prefixed" (since it contains a ``:``) and skips the
``ValueError`` raise
+ entirely -- the ``pytest.raises(ValueError, match="primary")`` block
would then fail
+ because no exception is raised (the stubbed ``infer_model`` would
resolve it instead).
+ """
+ registry.add("primary", extra={"model":
"us.anthropic.claude-opus-4-6-v1:0"})
+ hook = PydanticAIHook(llm_conn_id="primary")
+
+ with pytest.raises(ValueError, match="primary") as exc_info:
+ hook.get_conn()
+
+ assert "us.anthropic.claude-opus-4-6-v1:0" in str(exc_info.value)
+
+
@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_provider_class",
autospec=True)
+ def test_import_error_from_recognized_provider_counts_as_prefixed(self,
mock_infer_provider_class):
+ """A recognized provider name whose optional dependency isn't
installed still counts
+ as a platform prefix -- ``infer_provider_class`` raises
``ImportError`` (not
+ ``ValueError``) for a name it recognizes but can't import.
+
+ Mutation canary: changing the ``except ImportError`` branch to
``return False``
+ makes this assert ``True`` fail.
+ """
+ mock_infer_provider_class.side_effect = ImportError("Please install
the 'azure' extra")
+
+ assert _has_recognized_provider_prefix("azure:foo") is True
+
+
@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_provider_class",
autospec=True)
+ def test_value_error_from_unknown_provider_counts_as_bare(self,
mock_infer_provider_class):
+ """An unrecognized name raises ``ValueError`` and is treated as a bare
model name --
+ the counterpart to the ``ImportError`` case above, proving the two
exceptions are
+ told apart rather than both mapping to the same answer.
+
+ Mutation canary: changing the ``except ValueError`` branch to ``return
True``
+ makes this assert ``False`` fail.
+ """
+ mock_infer_provider_class.side_effect = ValueError("Unknown provider:
bogus")
+
+ assert _has_recognized_provider_prefix("bogus:foo") is False
+
+ @pytest.mark.parametrize(
+ ("prefix", "expected"),
+ [
+ ("google-vertex", True),
+ ("google-gla", True),
+ ("openi", True),
+ ("us.anthropic.claude-opus-4-6-v1", False), # Bedrock native id:
contains '.'
+ (
+ "azure",
+ True,
+ ), # shape-only check; caller only invokes this after ruling out
recognized prefixes
+ ],
+ )
+ def test_looks_like_unrecognized_provider_prefix(self, prefix, expected):
+ assert _looks_like_unrecognized_provider_prefix(prefix) is expected
+
+ def test_vertex_stale_gateway_prefix_warns_but_still_resolves(self,
registry, infer_model_stub, caplog):
+ registry.add(
+ "primary",
+ conn_type="pydanticai_vertex",
+ hook_class=PydanticAIVertexHook,
+ extra={"model": "google-vertex:gemini-2.0-flash", "api_key":
"some-key"},
+ )
+ hook = PydanticAIVertexHook(llm_conn_id="primary")
+
+ with caplog.at_level(logging.WARNING):
+ model = hook.get_conn()
+
+ assert model is
infer_model_stub.models["google-cloud:google-vertex:gemini-2.0-flash"]
+ assert any("google-vertex" in r.message and "typo" in r.message for r
in caplog.records)
+
+ def test_bedrock_native_id_does_not_warn(self, registry, infer_model_stub,
caplog):
+ """Mutation canary: dropping the '.' exclusion (or the shape check
entirely) makes this
+ assert fail -- the legitimate Bedrock id would start logging a warning
on every resolution.
+ """
+ registry.add(
+ "primary",
+ conn_type="pydanticai_bedrock",
+ hook_class=PydanticAIBedrockHook,
+ extra={"model": "us.anthropic.claude-opus-4-6-v1:0"},
+ )
+ hook = PydanticAIBedrockHook(llm_conn_id="primary")
+
+ with caplog.at_level(logging.WARNING):
+ hook.get_conn()
+
+ assert not any("typo" in r.message for r in caplog.records)
+
+ def test_generic_connection_typo_prefix_names_the_bad_segment(self,
registry, infer_model_stub):
+ """A colon-bearing name with an unrecognized prefix must name that
prefix, not tell the
+ user their already-set 'provider:model' string doesn't exist.
+ """
+ registry.add("primary", extra={"model": "openi:gpt-5"})
+ hook = PydanticAIHook(llm_conn_id="primary")
+
+ with pytest.raises(ValueError, match="primary") as exc_info:
+ hook.get_conn()
+
+ assert "openi" in str(exc_info.value)
+ assert "openi:gpt-5" in str(exc_info.value)
+
+
+class TestPydanticAIHookFallback:
+ def test_no_fallback_returns_the_bare_model(self, registry,
infer_model_stub):
+ """Without a chain the resolved model is not wrapped at all."""
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ hook = PydanticAIHook(llm_conn_id="primary")
+
+ assert hook.get_conn() is infer_model_stub.models["openai:gpt-5.6-sol"]
+
+ def test_param_builds_chain_in_order(self, registry, infer_model_stub):
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+ registry.add("third", extra={"model": "groq:llama-4"})
+
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["second", "third"])
+ model = hook.get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["openai:gpt-5.6-sol"],
+ infer_model_stub.models["anthropic:claude-opus-4-6"],
+ infer_model_stub.models["groq:llama-4"],
+ ]
+
+ @patch("airflow.providers.common.ai.hooks.pydantic_ai.FallbackModel",
autospec=True)
+ def test_chain_pins_fallback_on_to_model_api_error(self,
mock_fallback_model, registry, infer_model_stub):
+ """The chain must not drift with pydantic-ai's own ``fallback_on``
default.
+
+ Asserts our own call into ``FallbackModel`` -- not pydantic-ai's
dispatch logic, which
+ is a third party's private implementation detail -- so deleting the
``fallback_on=``
+ kwarg from the call site turns this red.
+ """
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["second"]).get_conn()
+
+ assert mock_fallback_model.mock_calls == [
+ call(
+ infer_model_stub.models["openai:gpt-5.6-sol"],
+ infer_model_stub.models["anthropic:claude-opus-4-6"],
+ fallback_on=(ModelAPIError,),
+ )
+ ]
+
+ def test_chain_from_connection_extra(self, registry, infer_model_stub):
+ """A deployment manager can configure failover without touching Dag
code."""
+ registry.add(
+ "primary",
+ extra={"model": "openai:gpt-5.6-sol", "fallback_conn_ids":
["second"]},
+ )
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ model = PydanticAIHook(llm_conn_id="primary").get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["openai:gpt-5.6-sol"],
+ infer_model_stub.models["anthropic:claude-opus-4-6"],
+ ]
+
+ def test_param_overrides_extra(self, registry, infer_model_stub):
+ registry.add(
+ "primary",
+ extra={"model": "openai:gpt-5.6-sol", "fallback_conn_ids":
["ignored"]},
+ )
+ registry.add("ignored", extra={"model": "groq:llama-4"})
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ model = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["second"]).get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models[1] is
infer_model_stub.models["anthropic:claude-opus-4-6"]
+
+ def test_empty_list_param_disables_the_extra_chain(self, registry,
infer_model_stub):
+ """``[]`` is an explicit opt-out, distinct from ``None`` meaning "read
the extra"."""
+ registry.add(
+ "primary",
+ extra={"model": "openai:gpt-5.6-sol", "fallback_conn_ids":
["second"]},
+ )
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ model = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=[]).get_conn()
+
+ assert model is infer_model_stub.models["openai:gpt-5.6-sol"]
+
+ def test_chain_can_span_providers(self, registry, infer_model_stub):
+ """Each connection resolves through its own hook class, so credentials
differ per hop."""
+ registry.add("primary", password="sk-openai", extra={"model":
"openai:gpt-5.6-sol"})
+ registry.add(
+ "bedrock_dr",
+ conn_type="pydanticai_bedrock",
+ hook_class=PydanticAIBedrockHook,
+ extra={
+ "model": "bedrock:us.anthropic.claude-opus-4-5",
+ "region_name": "us-east-1",
+ "aws_access_key_id": "AKIA-test",
+ "aws_secret_access_key": "secret",
+ },
+ )
+
+ with patch(
+
"airflow.providers.common.ai.hooks.pydantic_ai.infer_provider_class",
autospec=True
+ ) as mock_infer_provider_class:
+ mock_infer_provider_class.return_value =
MagicMock(return_value=MagicMock())
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["bedrock_dr"])
+ model = hook.get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["openai:gpt-5.6-sol"],
+
infer_model_stub.models["bedrock:us.anthropic.claude-opus-4-5"],
+ ]
+
+ # Each hop is built by its own hook's field mapping: the primary
from
+ # password/host, the Bedrock hop from its extra.
+ assert infer_model_stub.provider_kwargs_for("openai:gpt-5.6-sol",
mock_infer_provider_class) == {
+ "api_key": "sk-openai"
+ }
+ assert infer_model_stub.provider_kwargs_for(
+ "bedrock:us.anthropic.claude-opus-4-5",
mock_infer_provider_class
+ ) == {
+ "region_name": "us-east-1",
+ "aws_access_key_id": "AKIA-test",
+ "aws_secret_access_key": "secret",
+ }
+
+ def test_bare_model_id_forwarded_to_fallback_without_own_model(self,
registry, infer_model_stub):
+ """A bare ``model_id`` flows to a fallback with none, qualified with
*that* fallback's platform."""
+ registry.add("primary", conn_type="pydanticai_azure",
hook_class=PydanticAIAzureHook)
+ registry.add("bedrock_dr", conn_type="pydanticai_bedrock",
hook_class=PydanticAIBedrockHook)
+
+ hook = PydanticAIAzureHook(
+ llm_conn_id="primary", model_id="gpt-5-nano",
fallback_conn_ids=["bedrock_dr"]
+ )
+ model = hook.get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["azure:gpt-5-nano"],
+ infer_model_stub.models["bedrock:gpt-5-nano"],
+ ]
+
+ def
test_bare_connection_model_forwarded_to_fallback_without_own_model(self,
registry, infer_model_stub):
+ """A bare model from the primary's own ``extra`` forwards like a
``model_id`` argument.
+
+ This is the connection-driven shape the docs lead with: neither the
model nor the chain
+ is named in Dag code, so forwarding the constructor argument alone
never fires.
+
+ Mutation canary: forwarding ``self.model_id`` rather than the
primary's configured name
+ makes ``get_conn()`` raise "No model specified for connection
'bedrock_dr'" here, because
+ ``model_id`` is ``None`` in this shape -- failing before either
assertion is reached.
+ """
+ registry.add(
+ "primary",
+ conn_type="pydanticai_azure",
+ hook_class=PydanticAIAzureHook,
+ extra={"model": "gpt-5-nano", "fallback_conn_ids": ["bedrock_dr"]},
+ )
+ registry.add("bedrock_dr", conn_type="pydanticai_bedrock",
hook_class=PydanticAIBedrockHook)
+
+ model = PydanticAIAzureHook(llm_conn_id="primary").get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["azure:gpt-5-nano"],
+ infer_model_stub.models["bedrock:gpt-5-nano"],
+ ]
+
+ def test_fallback_own_model_overrides_forwarded(self, registry,
infer_model_stub):
+ """A fallback's own ``model`` extra wins over anything forwarded from
the primary."""
+ registry.add("primary", conn_type="pydanticai_azure",
hook_class=PydanticAIAzureHook)
+ registry.add(
+ "bedrock_dr",
+ conn_type="pydanticai_bedrock",
+ hook_class=PydanticAIBedrockHook,
+ extra={"model": "bedrock:us.anthropic.claude-opus-4-5"},
+ )
+
+ hook = PydanticAIAzureHook(
+ llm_conn_id="primary", model_id="gpt-5-nano",
fallback_conn_ids=["bedrock_dr"]
+ )
+ model = hook.get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models[1] is
infer_model_stub.models["bedrock:us.anthropic.claude-opus-4-5"]
+
+ def test_prefixed_model_id_not_forwarded_to_fallback(self, registry,
infer_model_stub):
+ """A prefixed ``model_id`` pins the primary's own platform and is
unusable on a fallback."""
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second") # no model of its own
+
+ hook = PydanticAIHook(llm_conn_id="primary", model_id="openai:gpt-5",
fallback_conn_ids=["second"])
+ with pytest.raises(ValueError, match="No model specified for
connection 'second'"):
+ hook.get_conn()
+
+ def test_bedrock_style_bare_model_id_not_forwarded_across_platforms(self,
registry, infer_model_stub):
+ """A primary's bare model id with an embedded ``:`` of its own is a
native vendor id, and a
+ native vendor id is never valid on another platform -- it must not be
forwarded there.
+
+ Mutation canary: dropping the ``":" not in forwarded_model_id or
forwarded_model_provider
+ == self.model_provider`` gate (reverting to just ``not
_has_recognized_provider_prefix``)
+ makes the Bedrock id forward to the Azure fallback and this raise
never fires --
+ ``pytest.raises`` would report no exception raised.
+ """
+ registry.add("primary", conn_type="pydanticai_bedrock",
hook_class=PydanticAIBedrockHook)
+ registry.add("azure_dr", conn_type="pydanticai_azure",
hook_class=PydanticAIAzureHook)
+
+ hook = PydanticAIBedrockHook(
+ llm_conn_id="primary",
+ model_id="us.anthropic.claude-opus-4-6-v1:0",
+ fallback_conn_ids=["azure_dr"],
+ )
+ with pytest.raises(ValueError, match="No model specified for
connection 'azure_dr'"):
+ hook.get_conn()
+
+ def test_bedrock_style_bare_model_id_forwarded_within_same_platform(self,
registry, infer_model_stub):
+ """The cross-platform gate must not block Bedrock-to-Bedrock
forwarding of the same shape."""
+ registry.add("primary", conn_type="pydanticai_bedrock",
hook_class=PydanticAIBedrockHook)
+ registry.add("bedrock_dr", conn_type="pydanticai_bedrock",
hook_class=PydanticAIBedrockHook)
+
+ hook = PydanticAIBedrockHook(
+ llm_conn_id="primary",
+ model_id="us.anthropic.claude-opus-4-6-v1:0",
+ fallback_conn_ids=["bedrock_dr"],
+ )
+ model = hook.get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+
infer_model_stub.models["bedrock:us.anthropic.claude-opus-4-6-v1:0"],
+
infer_model_stub.models["bedrock:us.anthropic.claude-opus-4-6-v1:0"],
+ ]
+
+ def test_forwarded_bare_name_error_attributes_to_primary(self, registry,
infer_model_stub):
+ """A bare name forwarded from the primary must not blame the fallback
for a name it never
+ configured -- the error must name the primary connection that actually
set it.
+
+ Mutation canary: dropping ``forwarded_from_conn_id`` threading
(reverting
+ ``_qualify_model_name`` to always report the current connection as the
source) makes this
+ assert fail: the message would say the name came from 'openai_dr'
itself, not 'primary'.
+ """
+ registry.add(
+ "primary",
+ conn_type="pydanticai_bedrock",
+ hook_class=PydanticAIBedrockHook,
+ extra={"model": "claude-opus-4-5", "fallback_conn_ids":
["openai_dr"]},
+ )
+ registry.add("openai_dr") # generic connection type, no
model_provider, no own model
+
+ hook = PydanticAIBedrockHook(llm_conn_id="primary")
+ with pytest.raises(ValueError, match="forwarded from primary
connection 'primary'") as exc_info:
+ hook.get_conn()
+
+ assert "openai_dr" in str(exc_info.value)
+ assert "claude-opus-4-5" in str(exc_info.value)
+
+ def test_non_pydanticai_fallback_raises(self, registry, infer_model_stub):
+ """``Connection.get_hook`` dispatches on conn_type alone and can
return anything."""
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("wrong_type", conn_type="langchain")
+ registry.hook_classes["wrong_type"] = MagicMock # type:
ignore[assignment]
+
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["wrong_type"])
+ with pytest.raises(ValueError, match="not a PydanticAIHook"):
+ hook.get_conn()
+
+ def test_nested_chain_raises(self, registry, infer_model_stub):
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add(
+ "second",
+ extra={"model": "anthropic:claude-opus-4-6", "fallback_conn_ids":
["third"]},
+ )
+ registry.add("third", extra={"model": "groq:llama-4"})
+
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["second"])
+ with pytest.raises(ValueError, match="second.*not resolved
recursively"):
+ hook.get_conn()
+
+ @pytest.mark.parametrize(
+ ("fallback_conn_ids", "match"),
+ [
+ pytest.param(["second", "second"], "more than once",
id="repeated-fallback"),
+ pytest.param(["primary"], "as one of its own fallbacks",
id="primary-repeated"),
+ ],
+ )
+ def test_duplicate_conn_id_raises(self, registry, infer_model_stub,
fallback_conn_ids, match):
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=fallback_conn_ids)
+ with pytest.raises(ValueError, match=match):
+ hook.get_conn()
+
+ @pytest.mark.parametrize(
+ "fallback_conn_ids",
+ [
+ pytest.param("second,third", id="comma-separated-string"),
+ pytest.param(["second", 3], id="non-string-entry"),
+ ],
+ )
+ def test_malformed_fallback_conn_ids_raises(self, registry,
infer_model_stub, fallback_conn_ids):
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=fallback_conn_ids)
+ with pytest.raises(ValueError, match="must be a list of connection
IDs"):
+ hook.get_conn()
+
+ @pytest.mark.parametrize(
+ "fallback_conn_ids",
+ [
+ pytest.param(["second", ""], id="trailing-blank-entry"),
+ pytest.param(["second", " "], id="trailing-whitespace-entry"),
+ pytest.param(["second", "\n"], id="trailing-newline-entry"),
+ ],
+ )
+ def test_blank_entries_are_dropped_from_param(self, registry,
infer_model_stub, fallback_conn_ids):
+ """The Fallback Connections field is a textarea split on newline whose
blur handler
+ only guards the all-blank case, so a trailing blank line is what most
saved chains
+ actually look like; it must resolve as a working chain, not raise.
+ """
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ model = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=fallback_conn_ids).get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["openai:gpt-5.6-sol"],
+ infer_model_stub.models["anthropic:claude-opus-4-6"],
+ ]
+
+ def test_blank_entries_are_dropped_from_extra(self, registry,
infer_model_stub):
+ """Same drop-blank behavior applies when the chain comes from
connection extra."""
+ registry.add(
+ "primary",
+ extra={"model": "openai:gpt-5.6-sol", "fallback_conn_ids":
["second", ""]},
+ )
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ model = PydanticAIHook(llm_conn_id="primary").get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["openai:gpt-5.6-sol"],
+ infer_model_stub.models["anthropic:claude-opus-4-6"],
+ ]
+
+ def test_blank_only_chain_behaves_like_no_chain(self, registry,
infer_model_stub):
+ """Dropping every entry must fall back to the bare model, matching the
``[]`` opt-out."""
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+
+ model = PydanticAIHook(llm_conn_id="primary", fallback_conn_ids=["", "
"]).get_conn()
+
+ assert model is infer_model_stub.models["openai:gpt-5.6-sol"]
+
+ def test_fallback_entries_are_stripped(self, registry, infer_model_stub):
+ """Whitespace around a kept entry must not leak into the connection
lookup."""
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ model = PydanticAIHook(llm_conn_id="primary", fallback_conn_ids=["
second "]).get_conn()
+
+ assert isinstance(model, FallbackModel)
+ assert model.models == [
+ infer_model_stub.models["openai:gpt-5.6-sol"],
+ infer_model_stub.models["anthropic:claude-opus-4-6"],
+ ]
+
+ @pytest.mark.parametrize(
+ "malformed",
+ [
+ pytest.param("", id="empty-string"),
+ pytest.param(0, id="zero"),
+ pytest.param({}, id="empty-dict"),
+ pytest.param(False, id="false"),
+ ],
+ )
+ def test_malformed_fallback_conn_ids_from_extra_raises(self, registry,
infer_model_stub, malformed):
+ """Falsy-but-not-``None`` extra values must not be silently treated as
"no chain".
+
+ ``None`` is the only value the schema allows to mean "no chain"; any
other falsy
+ value that made it into the extra is a misconfiguration and must fail
the same way
+ a malformed hook argument does, not disappear into an empty chain.
+ """
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol",
"fallback_conn_ids": malformed})
+
+ hook = PydanticAIHook(llm_conn_id="primary")
+ with pytest.raises(ValueError, match="must be a list of connection
IDs"):
+ hook.get_conn()
+
+ def test_fallback_without_a_model_names_the_connection(self, registry,
infer_model_stub):
+ """The error has to say which hop is misconfigured, not just that one
is."""
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second")
+
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["second"])
+ with pytest.raises(ValueError, match="No model specified for
connection 'second'"):
+ hook.get_conn()
+
+ def test_test_connection_validates_the_whole_chain(self, registry,
infer_model_stub):
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ registry.add("second")
+
+ success, message = PydanticAIHook(
+ llm_conn_id="primary", fallback_conn_ids=["second"]
+ ).test_connection()
+
+ assert success is False
+ assert "second" in message
+
+ def test_missing_fallback_conn_is_reported_by_name(self, registry,
infer_model_stub):
+ """
+ A typo'd fallback conn_id must surface Airflow's real not-found error,
by name.
+
+ The test double has to fail the same way
``BaseHook.get_connection``/``get_hook`` do
+ in production (``AirflowNotFoundException``, not a bare ``KeyError``)
for this edge
+ case to be exercised at all. Both exception types happen to embed the
conn_id in
+ their message, so asserting on the type -- not just the message -- is
what actually
+ pins the fixture to the real behavior.
+ """
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["typo_conn"])
+
+ with pytest.raises(AirflowNotFoundException, match="typo_conn"):
+ hook.get_conn()
+
+ success, message = hook.test_connection()
+ assert success is False
+ assert "typo_conn" in message
+
+ def test_chain_log_fires_before_a_mid_resolution_failure(self, registry,
infer_model_stub, caplog):
+ """The chain-resolution log has to appear even when a later hop fails
to resolve.
+
+ Logging it only after the whole chain resolves would make it absent in
exactly the
+ case where it earns its keep: a typo'd fallback conn_id raises before
that point, and
+ on the connection-driven path the docs recommend, the Dag never
mentions the bad
+ conn_id at all, so the task log is the only place it could show up.
+ """
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=["typo_conn"])
+
+ with caplog.at_level(logging.INFO):
+ with pytest.raises(AirflowNotFoundException, match="typo_conn"):
+ hook.get_conn()
+
+ assert any("primary -> typo_conn" in r.message for r in caplog.records)
+
+ @pytest.mark.parametrize(
+ "fallback_conn_ids",
+ [
+ pytest.param(None, id="no-chain"),
+ pytest.param(["", " "], id="blank-only-chain"),
+ ],
+ )
+ def test_no_fallback_chain_does_not_log(self, registry, infer_model_stub,
caplog, fallback_conn_ids):
+ """An empty chain -- whether unset or collapsed from blank-only
entries -- logs nothing."""
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol"})
+ hook = PydanticAIHook(llm_conn_id="primary",
fallback_conn_ids=fallback_conn_ids)
+
+ with caplog.at_level(logging.INFO):
+ hook.get_conn()
+
+ assert not any("fallback chain" in r.message.lower() for r in
caplog.records)
+
+ def
test_exhausted_chain_from_a_connection_raises_fallback_exception_group(self,
registry):
+ """
+ A fully exhausted chain built by ``get_conn()`` surfaces pydantic-ai's
own aggregate
+ exception, not the last provider's own exception type. A ``RetryRule``
written against a
+ provider-specific exception has to account for that (see
``docs/retry_policies.rst``).
+ """
+
+ def _always_fails(messages, info):
+ raise ModelAPIError("test-model", "provider outage")
+
+ registry.add("primary", extra={"model": "openai:gpt-5.6-sol",
"fallback_conn_ids": ["second"]})
+ registry.add("second", extra={"model": "anthropic:claude-opus-4-6"})
+
+ with patch(
+ "airflow.providers.common.ai.hooks.pydantic_ai.infer_model",
+ autospec=True,
+ side_effect=lambda model, **kwargs: FunctionModel(_always_fails),
+ ):
+ model = PydanticAIHook(llm_conn_id="primary").get_conn()
+
+ with pytest.raises(FallbackExceptionGroup):
+ Agent(model, instructions="classify").run_sync("hello")
+
+
+class TestPydanticAIHookFallbackConnectionFetchCount:
+ """
+ ``TestPydanticAIHookFallback`` above patches ``Connection.get_hook``
directly (see
+ ``_ConnRegistry.get_hook``), so it never runs the real dispatch that
+ ``_resolve_fallback_models`` goes through -- that is exactly the code path
a double
+ connection-fetch per fallback hop would hide in. This mocks only
``get_connection`` and
+ lets ``Connection.get_hook`` run for real, to pin how many times each
connection in a
+ fallback chain is actually fetched.
+ """
+
+ def test_fallback_chain_fetches_each_connection_once(self,
infer_model_stub):
+ conns = {
+ "primary": Connection(
+ conn_id="primary",
+ conn_type="pydanticai",
+ extra=json.dumps({"model": "openai:gpt-4",
"fallback_conn_ids": ["fb1", "fb2"]}),
+ ),
+ "fb1": Connection(
+ conn_id="fb1", conn_type="pydanticai",
extra=json.dumps({"model": "anthropic:claude-1"})
+ ),
+ "fb2": Connection(
+ conn_id="fb2", conn_type="pydanticai",
extra=json.dumps({"model": "anthropic:claude-2"})
+ ),
+ }
+
+ def _get_connection(conn_id: str) -> Connection:
+ try:
+ return conns[conn_id]
+ except KeyError:
+ raise AirflowNotFoundException(f"The conn_id `{conn_id}` isn't
defined") from None
+
+ with patch.object(
+ PydanticAIHook, "get_connection", side_effect=_get_connection
+ ) as mock_get_connection:
+ PydanticAIHook(llm_conn_id="primary").get_conn()
+
+ # 3-connection chain (primary + 2 fallbacks): 1 fetch each = 3 total.
Before the
+ # `_seed_connection` fix, each fallback paid 2 fetches (one inside
+ # `PydanticAIHook.get_hook`, discarded, plus one more the first time
the new hook's
+ # own `_get_conn_and_extra` ran) = 1 + 2*2 = 5.
+ assert mock_get_connection.call_count == 3
+
+
class TestPydanticAIHookCreateAgent:
@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_model",
autospec=True)
@patch("airflow.providers.common.ai.hooks.pydantic_ai.Agent",
autospec=True)
@@ -328,6 +1194,24 @@ class TestPydanticAIHookCreateAgent:
)
mock_agent_cls.assert_not_called()
+ def
test_create_agent_with_spec_file_raises_when_chain_configured_without_model(self):
+ """A connection-only chain cannot be wired into a spec-file agent
silently.
+
+ The spec file's own model can't be wrapped in a ``FallbackModel`` --
it is
+ resolved by pydantic-ai, not by this hook -- so a connection that
declares
+ ``fallback_conn_ids`` but no ``model`` must fail loudly here rather
than let the
+ chain quietly disappear.
+ """
+ hook = PydanticAIHook(llm_conn_id="test_conn")
+ conn = Connection(
+ conn_id="test_conn",
+ conn_type="pydanticai",
+ extra=json.dumps({"fallback_conn_ids": ["second"]}),
+ )
+ with patch.object(hook, "get_connection", autospec=True,
return_value=conn):
+ with pytest.raises(ValueError, match="fallback chain is configured
for 'test_conn' but no model"):
+ hook.create_agent(spec_file="/path/to/agent.yaml")
+
@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_model",
autospec=True)
@patch("airflow.providers.common.ai.hooks.pydantic_ai.Agent")
def test_create_agent_with_spec_file_path_object(self, mock_agent_cls,
mock_infer_model):
@@ -485,9 +1369,14 @@ class TestPydanticAIHookTestConnection:
@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_model",
autospec=True)
def test_failed_connection(self, mock_infer_model):
- mock_infer_model.side_effect = ValueError("Unknown provider
'badprovider'")
+ """A recognized provider prefix with an unresolvable model still
reaches ``infer_model``.
+
+ ``model_id`` uses a real provider name (``openai``) so qualification
passes it through
+ verbatim; the failure being tested here is pydantic-ai's own, not this
hook's.
+ """
+ mock_infer_model.side_effect = ValueError("Unknown model
'nonexistent-model'")
- hook = PydanticAIHook(llm_conn_id="test_conn",
model_id="badprovider:model")
+ hook = PydanticAIHook(llm_conn_id="test_conn",
model_id="openai:nonexistent-model")
conn = Connection(
conn_id="test_conn",
conn_type="pydanticai",
@@ -496,7 +1385,7 @@ class TestPydanticAIHookTestConnection:
success, message = hook.test_connection()
assert success is False
- assert "Unknown provider" in message
+ assert "Unknown model" in message
def test_failed_connection_no_model(self):
hook = PydanticAIHook(llm_conn_id="test_conn")
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 da8360f24c5..a5d49b11926 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
@@ -176,6 +176,7 @@ class TestAgentOperatorTemplateFields:
"prompt",
"llm_conn_id",
"model_id",
+ "fallback_conn_ids",
"system_prompt",
"agent_params",
"message_history",
@@ -380,7 +381,9 @@ class TestAgentOperatorExecute:
result = op.execute(context=_make_context())
assert result == "The answer is 42."
- mock_hook_cls.get_hook.assert_called_once_with("my_llm",
hook_params={"model_id": None})
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids": None}
+ )
mock_hook_cls.get_hook.return_value.create_agent.assert_called_once_with(
output_type=str, instructions="You are helpful."
)
@@ -570,7 +573,47 @@ class TestAgentOperatorExecute:
)
op.execute(context=MagicMock())
- mock_hook_cls.get_hook.assert_called_once_with("my_llm",
hook_params={"model_id": "openai:gpt-5"})
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": "openai:gpt-5",
"fallback_conn_ids": None}
+ )
+
+ @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
+ def test_execute_forwards_fallback_conn_ids_to_hook(self, mock_hook_cls,
make_mock_run_result):
+ """``fallback_conn_ids`` on the operator overrides the connection's
own extra field."""
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
_make_mock_agent(
+ "ok", make_mock_run_result
+ )
+
+ op = AgentOperator(
+ task_id="test",
+ prompt="test",
+ llm_conn_id="my_llm",
+ fallback_conn_ids=["conn_a", "conn_b"],
+ )
+ op.execute(context=MagicMock())
+
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids":
["conn_a", "conn_b"]}
+ )
+
+ @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
+ def test_execute_forwards_empty_fallback_conn_ids_to_hook(self,
mock_hook_cls, make_mock_run_result):
+ """An explicit ``[]`` disables a chain configured on the connection,
not just an override."""
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
_make_mock_agent(
+ "ok", make_mock_run_result
+ )
+
+ op = AgentOperator(
+ task_id="test",
+ prompt="test",
+ llm_conn_id="my_llm",
+ fallback_conn_ids=[],
+ )
+ op.execute(context=MagicMock())
+
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids": []}
+ )
@pytest.mark.skipif(
not AIRFLOW_V_3_1_PLUS, reason="Human in the loop is only compatible
with Airflow >= 3.1.0"
diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llm.py
b/providers/common/ai/tests/unit/common/ai/operators/test_llm.py
index aec57fc8a44..a56d181d951 100644
--- a/providers/common/ai/tests/unit/common/ai/operators/test_llm.py
+++ b/providers/common/ai/tests/unit/common/ai/operators/test_llm.py
@@ -83,7 +83,15 @@ def _build_priced_response(messages: list[ModelMessage],
info: AgentInfo) -> Mod
class TestLLMOperator:
def test_template_fields(self):
- expected = {"prompt", "llm_conn_id", "model_id", "system_prompt",
"agent_params", "usage_limits"}
+ expected = {
+ "prompt",
+ "llm_conn_id",
+ "model_id",
+ "fallback_conn_ids",
+ "system_prompt",
+ "agent_params",
+ "usage_limits",
+ }
assert set(LLMOperator.template_fields) == expected
@pytest.mark.parametrize(
@@ -121,7 +129,9 @@ class TestLLMOperator:
mock_hook_cls.get_hook.return_value.create_agent.assert_called_once_with(
output_type=str, instructions=""
)
- mock_hook_cls.get_hook.assert_called_once_with("my_llm",
hook_params={"model_id": None})
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids": None}
+ )
@patch("airflow.providers.common.ai.operators.llm.PydanticAIHook",
autospec=True)
def test_execute_forwards_usage_limits_to_run_sync(self, mock_hook_cls,
make_mock_run_result):
@@ -270,7 +280,9 @@ class TestLLMOperator:
assert isinstance(result, Entities)
assert result.names == ["Alice", "Bob"]
- mock_hook_cls.get_hook.assert_called_once_with("my_llm",
hook_params={"model_id": "openai:gpt-5"})
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": "openai:gpt-5",
"fallback_conn_ids": None}
+ )
mock_hook_cls.get_hook.return_value.create_agent.assert_called_once_with(
output_type=Entities,
instructions="You are an extractor.",
@@ -278,6 +290,44 @@ class TestLLMOperator:
model_settings={"temperature": 0.9},
)
+ @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook",
autospec=True)
+ def test_execute_forwards_fallback_conn_ids_to_hook(self, mock_hook_cls,
make_mock_run_result):
+ """``fallback_conn_ids`` on the operator overrides the connection's
own extra field."""
+ mock_agent = MagicMock(spec=["run_sync"])
+ mock_agent.run_sync.return_value = make_mock_run_result("ok")
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
mock_agent
+
+ op = LLMOperator(
+ task_id="test",
+ prompt="p",
+ llm_conn_id="my_llm",
+ fallback_conn_ids=["conn_a", "conn_b"],
+ )
+ op.execute(context=MagicMock())
+
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids":
["conn_a", "conn_b"]}
+ )
+
+ @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook",
autospec=True)
+ def test_execute_forwards_empty_fallback_conn_ids_to_hook(self,
mock_hook_cls, make_mock_run_result):
+ """An explicit ``[]`` disables a chain configured on the connection,
not just an override."""
+ mock_agent = MagicMock(spec=["run_sync"])
+ mock_agent.run_sync.return_value = make_mock_run_result("ok")
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
mock_agent
+
+ op = LLMOperator(
+ task_id="test",
+ prompt="p",
+ llm_conn_id="my_llm",
+ fallback_conn_ids=[],
+ )
+ op.execute(context=MagicMock())
+
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids": []}
+ )
+
def test_declares_output_type_for_deserialization(self):
"""Declares ``output_type`` so the worker-side DAG walk registers it
for deserialization.
diff --git
a/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
b/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
index ff511183b16..d5ff67fd881 100644
---
a/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
+++
b/providers/common/ai/tests/unit/common/ai/operators/test_llm_file_analysis.py
@@ -85,6 +85,7 @@ class TestLLMFileAnalysisOperator:
"prompt",
"llm_conn_id",
"model_id",
+ "fallback_conn_ids",
"system_prompt",
"agent_params",
"usage_limits",
diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_llm_sql.py
b/providers/common/ai/tests/unit/common/ai/operators/test_llm_sql.py
index b96d836c490..70211d55c26 100644
--- a/providers/common/ai/tests/unit/common/ai/operators/test_llm_sql.py
+++ b/providers/common/ai/tests/unit/common/ai/operators/test_llm_sql.py
@@ -196,6 +196,7 @@ class TestLLMSQLQueryOperator:
"prompt",
"llm_conn_id",
"model_id",
+ "fallback_conn_ids",
"system_prompt",
"agent_params",
"usage_limits",
@@ -205,6 +206,25 @@ class TestLLMSQLQueryOperator:
}
assert set(LLMSQLQueryOperator.template_fields) == expected
+ @patch("airflow.providers.common.ai.operators.llm.PydanticAIHook",
autospec=True)
+ def test_execute_forwards_fallback_conn_ids_to_hook(self, mock_hook_cls,
make_mock_run_result):
+ """``fallback_conn_ids`` is accepted without a per-subclass code
change and reaches the hook."""
+ mock_agent = _make_mock_agent("SELECT id FROM users",
make_mock_run_result)
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
mock_agent
+
+ op = LLMSQLQueryOperator(
+ task_id="test",
+ prompt="Get users",
+ llm_conn_id="my_llm",
+ schema_context="Table: users\nColumns: id INT",
+ fallback_conn_ids=["conn_a", "conn_b"],
+ )
+ op.execute(context=MagicMock())
+
+ mock_hook_cls.get_hook.assert_called_once_with(
+ "my_llm", hook_params={"model_id": None, "fallback_conn_ids":
["conn_a", "conn_b"]}
+ )
+
@patch("airflow.providers.common.ai.operators.llm.PydanticAIHook",
autospec=True)
def test_execute_with_schema_context(self, mock_hook_cls,
make_mock_run_result):
"""Operator uses schema_context and returns generated SQL."""
diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
b/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
index 06389edf9d1..baf5ea01b48 100644
--- a/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
+++ b/providers/common/ai/tests/unit/common/ai/utils/test_logging.py
@@ -21,11 +21,16 @@ from decimal import Decimal
from unittest.mock import MagicMock
from pydantic import BaseModel
+from pydantic_ai import Agent
+from pydantic_ai.exceptions import ModelAPIError
from pydantic_ai.messages import (
ModelResponse,
ModelResponsePart,
ToolCallPart,
)
+from pydantic_ai.models.fallback import FallbackModel
+from pydantic_ai.models.function import FunctionModel
+from pydantic_ai.models.test import TestModel
from airflow.providers.common.ai.toolsets.logging import LoggingToolset
from airflow.providers.common.ai.utils.logging import (
@@ -93,6 +98,30 @@ class TestLogRunSummary:
assert tool_line == "Tool call sequence: list_tables -> get_schema ->
query"
assert records[-1].message == "::endgroup::"
+ def test_names_the_model_that_served_a_failed_over_run(self):
+ """
+ After a failover the summary names the model that actually answered.
+
+ A silent failover still reporting the primary would hide the cost and
quality
+ change from whoever has to account for which model produced an output.
+ """
+
+ def _primary_is_down(messages, info):
+ raise ModelAPIError("openai:gpt-5", "provider outage")
+
+ agent = Agent(
+ FallbackModel(FunctionModel(_primary_is_down), TestModel()),
+ instructions="classify",
+ )
+ result = agent.run_sync("hello")
+
+ logger = MagicMock(spec=logging.Logger)
+ log_run_summary(logger, result)
+
+ summary_format, *summary_args = logger.info.call_args_list[0].args
+ assert "model=%s" in summary_format
+ assert summary_args[0] == "test"
+
def test_no_tools_skips_sequence_line(self, caplog):
logger = logging.getLogger("test.log_run_summary")
result = _make_mock_result(tool_names=None)