Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 8 additions & 8 deletions providers/common/ai/docs/connections/pydantic_ai_vertex.rst
Original file line number Diff line number Diff line change
Expand Up @@ -60,14 +60,14 @@ Force Vertex AI Mode
above already makes ``GoogleCloudProvider`` hard-code ``vertexai=True``
unconditionally when it builds its client.

.. important::
Leave this field unset. Setting it currently breaks the connection:
neither ``GoogleProvider`` nor ``GoogleCloudProvider`` accept a
``vertexai`` constructor argument, so the hook silently discards
every other field on this connection (project, location, service
account, API key) and falls back to resolving credentials from
environment variables only. If auth unexpectedly falls back to env
vars, check the task log for a "rejected kwargs" warning.
.. note::
This field is accepted for backward compatibility but has no effect:
it is never forwarded to the provider, and every other field on the
connection (project, location, service account, API key) is passed
through normally. Setting it logs a warning in the task log noting
that the field is ignored and that Vertex AI vs. Generative Language
API mode is selected via the model prefix (``google-cloud:`` vs.
``google:``) instead.

API Key
Google API key for Vertex AI Express Mode. Falls back to the
Expand Down
4 changes: 2 additions & 2 deletions providers/common/ai/provider.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -310,7 +310,7 @@ connection-types:
conn-fields:
model:
label: Model
description: "Google model identifier (e.g. google-vertex:gemini-2.0-flash)"
description: "Google model identifier (e.g. google-cloud:gemini-2.0-flash)"
schema:
type:
- string
Expand All @@ -331,7 +331,7 @@ connection-types:
- 'null'
vertexai:
label: Force Vertex AI Mode
description: "Force Vertex AI mode. Auto-detected when project/location/credentials are set."
description: "Ignored (kept for compatibility); mode is now selected via the Model field's prefix."
schema:
type:
- boolean
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -248,7 +248,7 @@ def get_provider_info():
"conn-fields": {
"model": {
"label": "Model",
"description": "Google model identifier (e.g. google-vertex:gemini-2.0-flash)",
"description": "Google model identifier (e.g. google-cloud:gemini-2.0-flash)",
"schema": {"type": ["string", "null"]},
},
"project": {
Expand All @@ -263,7 +263,7 @@ def get_provider_info():
},
"vertexai": {
"label": "Force Vertex AI Mode",
"description": "Force Vertex AI mode. Auto-detected when project/location/credentials are set.",
"description": "Ignored (kept for compatibility); mode is now selected via the Model field's prefix.",
"schema": {"type": ["boolean", "null"]},
},
"api_key": {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -459,21 +459,22 @@ class PydanticAIVertexHook(PydanticAIHook):
- **extra** JSON::

{
"model": "google-vertex:gemini-2.0-flash",
"model": "google-cloud:gemini-2.0-flash",
"project": "my-gcp-project",
"location": "us-central1",
"service_account_info": {...},
"vertexai": true,
}

Use ``"service_account_info"`` to embed the service-account JSON directly
(as an object, not a string path).

Set ``"vertexai": true`` to force Vertex AI mode when only ``api_key`` is
provided. Omit ``vertexai`` for the Generative Language API (GLA).
``"vertexai"`` is accepted for backward compatibility but has no effect:
pydantic-ai now selects Vertex AI vs. the Generative Language API from the
model prefix (``google-cloud:`` vs. ``google:``) rather than a
constructor flag, so there is nothing left for this field to control.

:param llm_conn_id: Airflow connection ID.
:param model_id: Model identifier, e.g. ``"google-vertex:gemini-2.0-flash"``.
:param model_id: Model identifier, e.g. ``"google-cloud:gemini-2.0-flash"``.
"""

conn_type = "pydanticai-vertex"
Expand All @@ -488,8 +489,8 @@ def get_ui_field_behaviour() -> dict[str, Any]:
"relabeling": {},
"placeholders": {
"extra": (
'{"model": "google-vertex:gemini-2.0-flash", '
'"project": "my-project", "location": "us-central1", "vertexai": true}'
'{"model": "google-cloud:gemini-2.0-flash", '
'"project": "my-project", "location": "us-central1"}'
" — add service_account_info (object) for SA auth;"
" omit both to use Application Default Credentials"
),
Expand All @@ -510,10 +511,18 @@ def _get_provider_kwargs(
if extra.get(_key):
kwargs[_key] = extra[_key]

# Optional vertexai bool flag (force Vertex AI mode for API-key auth).
_vertexai = extra.get("vertexai")
if _vertexai is not None:
kwargs["vertexai"] = bool(_vertexai)
# "vertexai" predates pydantic-ai splitting GoogleProvider (Generative Language API)
# from GoogleCloudProvider (Vertex AI, which hardcodes vertexai=True internally and
# accepts no such constructor kwarg) in pydantic/pydantic-ai#5336. Forwarding it would
# raise TypeError, which the base hook's `except TypeError` in get_conn() would then
# swallow by falling back to env-var auth with *all* other kwargs discarded — silently
# authenticating as the wrong identity. Accept the field for backward compatibility but
# never forward it: which API is used is now controlled by the model prefix.
if extra.get("vertexai") is not None:
self.log.warning(
"The 'vertexai' connection field is ignored; Vertex AI vs. Generative Language "
"API mode is now selected via the model prefix ('google-cloud:' vs. 'google:')."
)

# Service-account credentials — loaded lazily to avoid importing
# google-auth on non-Vertex code paths (optional heavy dependency).
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -743,15 +743,31 @@ def test_get_provider_kwargs_api_key_gla_mode(self):
)
assert result["api_key"] == "gla-key"

def test_get_provider_kwargs_vertexai_flag(self):
"""vertexai bool is forwarded and coerced to bool."""
@pytest.mark.parametrize("vertexai_value", [True, False])
def test_get_provider_kwargs_vertexai_flag_is_not_forwarded(self, vertexai_value):
"""The ``vertexai`` extra field must never reach the provider constructor.

Neither ``GoogleProvider`` nor ``GoogleCloudProvider`` in current pydantic-ai
accept a ``vertexai`` kwarg (pydantic/pydantic-ai#5336 hardcoded it inside
``GoogleCloudProvider`` instead). Forwarding it raises ``TypeError``, which the
base hook's fallback then swallows by dropping every other kwarg -- silently
re-resolving credentials from the environment. Regression test for that bug.
"""
hook = PydanticAIVertexHook.__new__(PydanticAIVertexHook)
result = hook._get_provider_kwargs(
None,
None,
{"model": "google-vertex:gemini-2.0-flash", "api_key": "key", "vertexai": True},
{
"model": "google-vertex:gemini-2.0-flash",
"project": "my-project",
"location": "us-central1",
"vertexai": vertexai_value,
},
)
assert result["vertexai"] is True
assert "vertexai" not in result
# The other credential kwargs must still go through untouched.
assert result["project"] == "my-project"
assert result["location"] == "us-central1"

def test_get_provider_kwargs_service_account_info_loads_credentials(self):
"""service_account_info dict is loaded into a Credentials object."""
Expand Down Expand Up @@ -833,3 +849,67 @@ def test_get_conn_uses_explicit_project(self, mock_infer_provider_class, mock_in
factory = mock_infer_model.call_args[1]["provider_factory"]
factory("google-vertex")
mock_provider_cls.assert_called_with(project="my-project", location="europe-west4")

@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_model", autospec=True)
@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_provider", autospec=True)
@patch("airflow.providers.common.ai.hooks.pydantic_ai.infer_provider_class", autospec=True)
def test_get_conn_vertexai_flag_does_not_trigger_typeerror_fallback(
self, mock_infer_provider_class, mock_infer_provider, mock_infer_model
):
"""Setting ``vertexai`` must not push ``get_conn`` onto the ``except TypeError``
fallback path, which would silently discard project/location/credentials.

The stand-in below has the exact keyword-only signature of the real
``GoogleCloudProvider.__init__`` (verified against the installed pydantic-ai) so
it raises ``TypeError`` on an unexpected ``vertexai`` kwarg exactly like the real
class would -- the real class itself needs the optional ``google-genai``
dependency, which isn't part of this provider's test environment.
"""

class FakeGoogleCloudProvider:
def __init__(
self,
*,
api_key=None,
credentials=None,
project=None,
location=None,
client=None,
http_client=None,
base_url=None,
retry_options=None,
):
self.kwargs = {
"api_key": api_key,
"credentials": credentials,
"project": project,
"location": location,
}

mock_infer_model.return_value = MagicMock(spec=Model)
mock_infer_provider_class.return_value = FakeGoogleCloudProvider

hook = PydanticAIVertexHook(llm_conn_id="vertex_test")
conn = Connection(
conn_id="vertex_test",
conn_type="pydanticai-vertex",
extra=json.dumps(
{
"model": "google-cloud:gemini-2.0-flash",
"project": "my-project",
"location": "us-central1",
"vertexai": True,
}
),
)
with patch.object(hook, "get_connection", return_value=conn):
hook.get_conn()

factory = mock_infer_model.call_args[1]["provider_factory"]
provider = factory("google-cloud")

assert isinstance(provider, FakeGoogleCloudProvider)
assert provider.kwargs["project"] == "my-project"
assert provider.kwargs["location"] == "us-central1"
# The TypeError fallback must never have been reached.
mock_infer_provider.assert_not_called()