Lee-W commented on code in PR #71437:
URL: https://github.com/apache/airflow/pull/71437#discussion_r4102698081
##########
providers/common/ai/docs/connections/pydantic_ai.rst:
##########
@@ -54,14 +56,37 @@ Model
The model can also be overridden at the hook/operator level via the
``model_id`` parameter.
+Embedding Model
+ The embedding model identifier in ``provider:model`` format. This field
+ appears as a dedicated input in the connection form and stores its value in
+ ``extra["embed_model"]``.
+
+ Example: ``openai:text-embedding-3-small``
+
+ The embedding model and connection can also be overridden at the hook level
+ via the ``embed_model_id`` and ``embed_conn_id`` parameters.
+
+ When the LLM and embedding model use different providers, configure a
separate
+ ``embed_conn_id`` if the shared connection produces explicit configuration
for the
+ embedding provider. The providers can share a connection when no embedding
provider
+ configuration can be built from it; pydantic-ai then resolves the
embedding provider
+ independently, for example from environment variables, even if the LLM
uses configuration
+ from the connection. Equivalent OpenAI and Azure chat/response prefixes
can share their
+ provider's connection. Local ``sentence-transformers:`` embeddings can
also share the
+ LLM connection because they do not use provider credentials.
+
API Key (Password field)
- The API key for your LLM provider. Required for API-key-based providers
+ The API key for your model provider. Required for API-key-based providers
(OpenAI, Anthropic, Groq, Mistral). Leave empty for providers using
environment-based auth (Bedrock via ``AWS_PROFILE``, Vertex via
``GOOGLE_APPLICATION_CREDENTIALS``).
+ For Bedrock and Google models, provider-specific values from Extra are used
+ instead. A populated Password or Host is ignored for those model prefixes;
+ the hook emits a warning identifying the replacement Extra fields.
+
Review Comment:
```suggestion
For Bedrock and Google models, provider-specific values from "Extra"
field are used
instead. A populated value from "Password" or "Host" field is ignored
for those model prefixes;
the hook emits a warning identifying the replacement "Extra" fields.
```
##########
providers/common/ai/docs/hooks/pydantic_ai.rst:
##########
@@ -57,7 +58,46 @@ The model can be specified at three levels (highest priority
first):
# Override with a specific model
hook = PydanticAIHook(llm_conn_id="my_llm",
model_id="anthropic:claude-sonnet-5")
-Structured output
+Embedding Models
+----------------
+
+Set ``embed_model_id`` on the hook or ``embed_model`` in the connection's
extra JSON,
+then call ``get_embedder()``. ``embed_conn_id`` defaults to ``llm_conn_id``.
Different
+LLM and embedding providers require separate connections when the shared
connection
+produces explicit configuration for the embedding provider, so those values
cannot be
+reused for the wrong provider. They can share a connection when no embedding
provider
+configuration can be built from it; pydantic-ai then resolves the embedding
provider
+independently, for example from environment variables, even if the LLM uses
configuration
+from the connection. Equivalent OpenAI and Azure chat/response prefixes can
share their
+provider's embedding connection. Local
+``sentence-transformers:`` embeddings can also share the LLM connection
because they
+do not use provider credentials. The resolved ``Embedder`` is cached on the
hook instance.
+
+.. code-block:: python
+
+ hook = PydanticAIHook(
+ llm_conn_id="my_llm",
+ embed_conn_id="my_embeddings",
+ embed_model_id="openai:text-embedding-3-small",
+ )
+ embedder = hook.get_embedder()
+ result = embedder.embed_query_sync("Apache Airflow orchestrates
workflows.")
+ embedding = result.embeddings[0]
+
+Keyword arguments accepted by pydantic-ai's `Embedder constructor
+<https://ai.pydantic.dev/api/embeddings/#pydantic_ai.embeddings.Embedder.__init__>`__
+can be passed directly to ``get_embedder()``. These currently include
``settings`` and ``instrument``. Caller-supplied ``instrument`` takes
Review Comment:
```suggestion
can be passed directly to ``get_embedder()``, such as ``settings`` and
``instrument``. Caller-supplied ``instrument`` takes
```
##########
providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py:
##########
@@ -505,9 +604,57 @@ def _resolve_fallback_models(self) -> list[Model]:
return models
+ def get_embedder(self, **embedder_kwargs: Any) -> Embedder:
Review Comment:
The semantic meaning of `get_` feels unintuitive here. Since we're providing
`embedder_kwargs`, it's essentially reconstructing the embedder.
##########
providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py:
##########
@@ -151,12 +182,16 @@ def __init__(
# 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
+ self.embed_conn_id = embed_conn_id if embed_conn_id is not None else
self.llm_conn_id
+ self.embed_model_id = embed_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
Review Comment:
```suggestion
# ``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.embed_conn_id = embed_conn_id if embed_conn_id is not None else
self.llm_conn_id
self.embed_model_id = embed_model_id
```
nit: a bit easier to read if we keep the order consistent
##########
providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py:
##########
@@ -115,25 +128,47 @@ class TestPydanticAIHookInit:
def test_default_conn_id(self):
hook = PydanticAIHook()
assert hook.llm_conn_id == "pydanticai_default"
+ assert hook.embed_conn_id == "pydanticai_default"
assert hook.model_id is None
+ assert hook.embed_model_id is None
def test_custom_conn_id(self):
- hook = PydanticAIHook(llm_conn_id="my_llm",
model_id="openai:gpt-5.6-sol")
+ hook = PydanticAIHook(
+ llm_conn_id="my_llm",
+ model_id="openai:gpt-5.6-sol",
+ embed_model_id="openai:text-embedding-3-small",
+ embed_conn_id="my_embeddings",
+ )
assert hook.llm_conn_id == "my_llm"
+ assert hook.embed_conn_id == "my_embeddings"
assert hook.model_id == "openai:gpt-5.6-sol"
+ assert hook.embed_model_id == "openai:text-embedding-3-small"
def test_azure_hook_uses_own_default_conn_name(self):
Review Comment:
Let's parameterize these tests
##########
providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py:
##########
@@ -203,12 +238,14 @@ def _get_provider_kwargs(
kwargs["base_url"] = base_url
return kwargs
- def _get_conn_and_extra(self) -> tuple[Connection, dict[str, Any]]:
+ def _get_conn_and_extra(self, conn_id) -> tuple[Connection, dict[str,
Any]]:
Review Comment:
```suggestion
def _get_conn_and_extra(self, conn_id: str | None) -> tuple[Connection,
dict[str, Any]]:
```
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]