This is an automated email from the ASF dual-hosted git repository.
kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new ede78e095b6 Add embedding kwargs to common AI hooks and operators
(#72002)
ede78e095b6 is described below
commit ede78e095b6dd72a2db707606ab5c927c4f828c8
Author: Jeff(Wei-Hao) Lu <[email protected]>
AuthorDate: Thu Oct 1 03:25:19 2026 +0800
Add embedding kwargs to common AI hooks and operators (#72002)
---
providers/common/ai/docs/hooks/langchain.rst | 20 +++++-
providers/common/ai/docs/hooks/llamaindex.rst | 12 ++++
.../ai/docs/operators/llamaindex_embedding.rst | 26 ++++++--
.../ai/docs/operators/llamaindex_retrieval.rst | 18 +++++-
.../airflow/providers/common/ai/hooks/langchain.py | 30 ++++++++-
.../providers/common/ai/hooks/llamaindex.py | 25 +++++++-
.../common/ai/operators/llamaindex_embedding.py | 13 ++++
.../common/ai/operators/llamaindex_retrieval.py | 14 +++++
.../tests/unit/common/ai/hooks/test_langchain.py | 66 ++++++++++++++++++++
.../tests/unit/common/ai/hooks/test_llamaindex.py | 72 ++++++++++++++++++++++
.../ai/operators/test_llamaindex_embedding.py | 15 ++++-
.../ai/operators/test_llamaindex_retrieval.py | 15 ++++-
12 files changed, 313 insertions(+), 13 deletions(-)
diff --git a/providers/common/ai/docs/hooks/langchain.rst
b/providers/common/ai/docs/hooks/langchain.rst
index b141b36e881..1783dfa6ca4 100644
--- a/providers/common/ai/docs/hooks/langchain.rst
+++ b/providers/common/ai/docs/hooks/langchain.rst
@@ -162,7 +162,25 @@ Parameters
- ``None`` (falls back to ``extra["embed_model"]`` on the connection)
- Embedding model identifier in ``provider:name`` form, e.g.
``openai:text-embedding-3-small``. Only required when calling
- ``get_embedding_model()``.
+ ``get_embedding_model()``. When ``embedding_kwargs`` supplies
+ ``provider`` explicitly, use a model name without the provider prefix.
+ * - ``embedding_kwargs``
+ - ``None``
+ - Additional keyword arguments passed to the embedding model constructor,
+ for example ``{"dimensions": 128}``. Values are forwarded without
+ filtering and can override hook-provided settings, including the
endpoint
+ and credentials. In particular, ``provider`` takes precedence over the
+ provider inferred from ``embed_model``. When ``provider`` is set,
+ LangChain treats the entire ``embed_model`` value as the model name
rather
+ than parsing a ``provider:name`` identifier. The hook logs a warning
when
+ both forms are supplied. Connection ``api_key`` and ``base_url`` values
+ take precedence over the same top-level keys, but the underlying
integration
+ may accept alternative or nested options that take precedence. Only pass
+ trusted values.
+
+.. seealso::
+ `langchain.embeddings.init_embeddings
<https://reference.langchain.com/python/langchain/embeddings/base/init_embeddings>`__
+ for valid ``embedding_kwargs`` keys.
Dependencies
------------
diff --git a/providers/common/ai/docs/hooks/llamaindex.rst
b/providers/common/ai/docs/hooks/llamaindex.rst
index 44dd4e2335c..4f35f4b10c5 100644
--- a/providers/common/ai/docs/hooks/llamaindex.rst
+++ b/providers/common/ai/docs/hooks/llamaindex.rst
@@ -115,10 +115,22 @@ Parameters
* - ``embed_model``
- ``None`` (falls back to ``extra["embed_model"]``)
- Embedding model name, e.g. ``text-embedding-3-small``.
+ * - ``embedding_kwargs``
+ - ``None``
+ - Additional keyword arguments passed to ``OpenAIEmbedding``, for example
+ ``{"dimensions": 128}``. Values are forwarded without filtering.
+ Connection ``api_key`` and ``api_base`` values take precedence at the
top
+ level, but nested options supported by the underlying library can
override
+ hook-provided request values, including credentials, the model, and the
+ input. Only pass trusted values.
* - ``llm_model``
- ``None`` (falls back to ``extra["llm_model"]``)
- LLM model name, e.g. ``gpt-5``. Required when calling ``get_llm()``.
+.. seealso::
+ `llama_index.embeddings.openai.OpenAIEmbedding
<https://developers.llamaindex.ai/python/framework-api-reference/embeddings/openai/>`__
+ for valid ``embedding_kwargs`` keys.
+
Dependencies
------------
diff --git a/providers/common/ai/docs/operators/llamaindex_embedding.rst
b/providers/common/ai/docs/operators/llamaindex_embedding.rst
index d3a101c3f81..a3ba933dfaa 100644
--- a/providers/common/ai/docs/operators/llamaindex_embedding.rst
+++ b/providers/common/ai/docs/operators/llamaindex_embedding.rst
@@ -91,22 +91,38 @@ Parameters
binding ``loader.output`` resolves to the native list before
execute.
* - ``embed_model``
- - String model name OR pre-built ``BaseEmbedding`` instance.
+ - String model name OR pre-built ``BaseEmbedding`` instance. Templated.
* - ``llm_conn_id``
- Airflow connection ID used when ``embed_model`` is a string. Falls
back to ``LlamaIndexHook.default_conn_name`` (``llamaindex_default``)
- when ``None``.
+ when ``None``. Templated.
* - ``embed_conn_id``
- Optional separate connection ID for the embedding provider. Falls
- back to ``llm_conn_id`` when ``None``.
+ back to ``llm_conn_id`` when ``None``. Templated.
+ * - ``embedding_kwargs``
+ - Additional keyword arguments passed to the embedding model constructor
+ when ``embed_model`` is a string or omitted, for example
+ ``{"dimensions": 128}``. Templated, so binding an upstream task's output
+ resolves to the native dictionary before execute and preserves typed
+ values such as integer ``dimensions``. When persisting an index, record
+ and reuse shape-affecting values such as ``dimensions`` in the retrieval
+ operator's ``embedding_kwargs``. Values are forwarded without filtering.
+ Connection credentials take precedence at the top level, but nested
+ options supported by the underlying library can override hook-provided
+ request values, including credentials, the model, and the input. Only
pass
+ trusted values.
* - ``chunk_size``
- Sentence-splitter chunk size (default 512).
* - ``chunk_overlap``
- Overlap between chunks (default 50).
* - ``persist_dir``
- - Local path or storage URI to persist the LlamaIndex index.
+ - Local path or storage URI to persist the LlamaIndex index. Templated.
* - ``persist_conn_id``
- - Cloud credentials connection ID for ``persist_dir`` URIs.
+ - Cloud credentials connection ID for ``persist_dir`` URIs. Templated.
+
+.. seealso::
+ `llama_index.embeddings.openai.OpenAIEmbedding
<https://developers.llamaindex.ai/python/framework-api-reference/embeddings/openai/>`__
+ for valid ``embedding_kwargs`` keys.
Output
------
diff --git a/providers/common/ai/docs/operators/llamaindex_retrieval.rst
b/providers/common/ai/docs/operators/llamaindex_retrieval.rst
index 1a0b254023b..399edbce13c 100644
--- a/providers/common/ai/docs/operators/llamaindex_retrieval.rst
+++ b/providers/common/ai/docs/operators/llamaindex_retrieval.rst
@@ -88,13 +88,27 @@ Parameters
* - ``llm_conn_id``
- Airflow connection ID used when ``embed_model`` is a string. Falls
back to ``LlamaIndexHook.default_conn_name`` (``llamaindex_default``)
- when ``None``.
+ when ``None``. Templated.
* - ``embed_conn_id``
- Optional separate connection ID for the embedding provider. Falls
- back to ``llm_conn_id`` when ``None``.
+ back to ``llm_conn_id`` when ``None``. Templated.
+ * - ``embedding_kwargs``
+ - Additional keyword arguments passed to the embedding model constructor
+ when ``embed_model`` is a string or omitted. Options such as
+ ``dimensions`` must match those used to build the index. Templated, so
+ binding an upstream task's output resolves to the native dictionary
+ before execute and preserves typed values such as integer
``dimensions``.
+ Values are forwarded without filtering. Connection credentials take
+ precedence at the top level, but nested options supported by the
underlying
+ library can override hook-provided request values, including
credentials,
+ the model, and the input. Only pass trusted values.
* - ``top_k``
- Number of top similarity results to return (default 5).
+.. seealso::
+ `llama_index.embeddings.openai.OpenAIEmbedding
<https://developers.llamaindex.ai/python/framework-api-reference/embeddings/openai/>`__
+ for valid ``embedding_kwargs`` keys.
+
Output
------
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/hooks/langchain.py
b/providers/common/ai/src/airflow/providers/common/ai/hooks/langchain.py
index 55b779bba0f..1d6122359a1 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/hooks/langchain.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/hooks/langchain.py
@@ -71,7 +71,20 @@ class LangChainHook(BaseHook):
Overrides ``extra["model"]`` on the connection.
:param embed_model: Embedding model identifier in ``provider:name`` format
(e.g. ``"openai:text-embedding-3-small"``). Overrides
- ``extra["embed_model"]`` on the connection.
+ ``extra["embed_model"]`` on the connection. When ``embedding_kwargs``
+ supplies ``provider`` explicitly, use a model name without the provider
+ prefix.
+ :param embedding_kwargs: Additional keyword arguments to pass to the
embedding
+ model constructor without filtering. Values can override hook-provided
+ settings, including the endpoint and credentials. In particular,
+ ``provider`` takes precedence over the provider inferred from
+ ``embed_model``. When ``provider`` is set, LangChain treats the entire
+ ``embed_model`` value as the model name rather than parsing a
+ ``provider:name`` identifier. The hook logs a warning when both forms
are
+ supplied. Connection ``api_key`` and ``base_url`` values take
precedence
+ over the same top-level keys, but the underlying integration may accept
+ alternative or nested options that take precedence. Only pass trusted
+ values.
"""
conn_name_attr = "llm_conn_id"
@@ -85,6 +98,8 @@ class LangChainHook(BaseHook):
embed_conn_id: str | None = None,
llm_model: str | None = None,
embed_model: str | None = None,
+ *,
+ embedding_kwargs: dict[str, Any] | None = None,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
@@ -96,6 +111,7 @@ class LangChainHook(BaseHook):
self.embed_conn_id = embed_conn_id if embed_conn_id is not None else
self.llm_conn_id
self.llm_model = llm_model
self.embed_model = embed_model
+ self.embedding_kwargs = embedding_kwargs or {}
@staticmethod
def get_ui_field_behaviour() -> dict[str, Any]:
@@ -174,7 +190,17 @@ class LangChainHook(BaseHook):
extra_key="embed_model",
kind="embedding",
)
- return init_embeddings(model_id, **self._connection_kwargs(conn))
+ connection_kwargs = self._connection_kwargs(conn)
+ overridden_keys = sorted(self.embedding_kwargs.keys() &
connection_kwargs.keys())
+ if overridden_keys:
+ self.log.warning("Connection parameters override embedding_kwargs
values: %s", overridden_keys)
+ if self.embedding_kwargs.get("provider") is not None and ":" in
model_id:
+ self.log.warning(
+ "embedding_kwargs['provider'] takes precedence over the
provider prefix in embed_model; "
+ "pass an unprefixed model name"
+ )
+ kwargs = {**self.embedding_kwargs, **connection_kwargs}
+ return init_embeddings(model_id, **kwargs)
def test_connection(self) -> tuple[bool, str]:
"""
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/hooks/llamaindex.py
b/providers/common/ai/src/airflow/providers/common/ai/hooks/llamaindex.py
index f0b45370b5b..6a3b01f9edd 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/hooks/llamaindex.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/hooks/llamaindex.py
@@ -18,6 +18,7 @@
from __future__ import annotations
+import inspect
from typing import TYPE_CHECKING, Any
from airflow.providers.common.compat.sdk import (
@@ -98,6 +99,11 @@ class LlamaIndexHook(BaseHook):
:param llm_model: LLM model name (e.g. ``"gpt-5"``). Overrides
``extra["llm_model"]`` on the connection. Required when calling
:meth:`get_llm`.
+ :param embedding_kwargs: Additional keyword arguments to pass to the
embedding
+ model constructor without filtering. Connection ``api_key`` and
``api_base``
+ values take precedence at the top level, but nested options supported
by the
+ underlying library can override hook-provided request values, including
+ credentials, the model, and the input. Only pass trusted values.
"""
conn_name_attr = "llm_conn_id"
@@ -111,6 +117,8 @@ class LlamaIndexHook(BaseHook):
embed_conn_id: str | None = None,
embed_model: str | None = None,
llm_model: str | None = None,
+ *,
+ embedding_kwargs: dict[str, Any] | None = None,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
@@ -119,6 +127,7 @@ class LlamaIndexHook(BaseHook):
self.llm_conn_id = llm_conn_id if llm_conn_id is not None else
self.default_conn_name
self.embed_conn_id = embed_conn_id if embed_conn_id is not None else
self.llm_conn_id
self.embed_model = embed_model
+ self.embedding_kwargs = embedding_kwargs or {}
self.llm_model = llm_model
@staticmethod
@@ -184,7 +193,21 @@ class LlamaIndexHook(BaseHook):
extra_key="embed_model",
kind="embedding",
)
- return OpenAIEmbedding(model=model_id, **self._connection_kwargs(conn))
+ connection_kwargs = self._connection_kwargs(conn)
+ overridden_keys = sorted(self.embedding_kwargs.keys() &
connection_kwargs.keys())
+ if overridden_keys:
+ self.log.warning("Connection parameters override embedding_kwargs
values: %s", overridden_keys)
+ kwargs = {**self.embedding_kwargs, **connection_kwargs}
+ supported_kwargs = {
+ name
+ for name, parameter in
inspect.signature(OpenAIEmbedding.__init__).parameters.items()
+ if name != "self"
+ and parameter.kind not in {inspect.Parameter.VAR_POSITIONAL,
inspect.Parameter.VAR_KEYWORD}
+ } | set(OpenAIEmbedding.model_fields)
+ unsupported_keys = sorted(self.embedding_kwargs.keys() -
supported_kwargs)
+ if unsupported_keys:
+ self.log.warning("OpenAIEmbedding ignores unsupported
embedding_kwargs: %s", unsupported_keys)
+ return OpenAIEmbedding(model=model_id, **kwargs)
def get_llm(self) -> LLM:
"""
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py
index bfc8001091f..c7003db1e33 100644
---
a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py
+++
b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_embedding.py
@@ -74,6 +74,11 @@ class LlamaIndexEmbeddingOperator(BaseOperator):
back to :attr:`LlamaIndexHook.default_conn_name` when ``None``.
:param embed_conn_id: Optional separate Airflow connection ID for the
embedding provider. Falls back to ``llm_conn_id`` when ``None``.
+ :param embedding_kwargs: Additional keyword arguments passed to the
embedding
+ model constructor without filtering when ``embed_model`` is a string or
+ omitted. Nested options supported by the underlying library can
override
+ hook-provided request values, including credentials, the model, and the
+ input. Only pass trusted values.
:param chunk_size: Chunk size for the sentence splitter.
:param chunk_overlap: Overlap between chunks.
:param persist_dir: Optional path to persist the index. Accepts local
@@ -88,6 +93,7 @@ class LlamaIndexEmbeddingOperator(BaseOperator):
"embed_model",
"llm_conn_id",
"embed_conn_id",
+ "embedding_kwargs",
"persist_dir",
"persist_conn_id",
)
@@ -99,6 +105,7 @@ class LlamaIndexEmbeddingOperator(BaseOperator):
embed_model: str | BaseEmbedding | None = None,
llm_conn_id: str | None = None,
embed_conn_id: str | None = None,
+ embedding_kwargs: dict[str, Any] | None = None,
chunk_size: int = 512,
chunk_overlap: int = 50,
persist_dir: str | None = None,
@@ -110,6 +117,7 @@ class LlamaIndexEmbeddingOperator(BaseOperator):
self.embed_model = embed_model
self.llm_conn_id = llm_conn_id
self.embed_conn_id = embed_conn_id
+ self.embedding_kwargs = embedding_kwargs or {}
self.chunk_size = chunk_size
self.chunk_overlap = chunk_overlap
self.persist_dir = persist_dir
@@ -195,6 +203,7 @@ class LlamaIndexEmbeddingOperator(BaseOperator):
llm_conn_id=self.llm_conn_id,
embed_conn_id=self.embed_conn_id,
embed_model=self.embed_model,
+ embedding_kwargs=self.embedding_kwargs,
).get_embedding_model()
# ``BaseEmbedding`` always exposes these two methods (see
@@ -205,6 +214,10 @@ class LlamaIndexEmbeddingOperator(BaseOperator):
if hasattr(self.embed_model, "get_text_embedding_batch") and hasattr(
self.embed_model, "_get_query_embedding"
):
+ if self.embedding_kwargs:
+ self.log.warning(
+ "embedding_kwargs is ignored when embed_model is a
pre-built embedding model"
+ )
return self.embed_model
raise TypeError(
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py
index 67f8fa9bfa3..234a3bfdd4d 100644
---
a/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py
+++
b/providers/common/ai/src/airflow/providers/common/ai/operators/llamaindex_retrieval.py
@@ -77,6 +77,12 @@ class LlamaIndexRetrievalOperator(BaseOperator):
Used only when ``embed_model`` is a string (or omitted entirely).
:param embed_conn_id: Optional separate Airflow connection ID for the
embedding provider. Falls back to ``llm_conn_id`` when ``None``.
+ :param embedding_kwargs: Additional keyword arguments passed to the
embedding
+ model constructor without filtering when ``embed_model`` is a string or
+ omitted. Options that affect vector dimensions must match those used to
+ build the index. Nested options supported by the underlying library can
+ override hook-provided request values, including credentials, the
model,
+ and the input. Only pass trusted values.
:param top_k: Number of top results to retrieve.
"""
@@ -87,6 +93,7 @@ class LlamaIndexRetrievalOperator(BaseOperator):
"embed_model",
"llm_conn_id",
"embed_conn_id",
+ "embedding_kwargs",
)
def __init__(
@@ -98,6 +105,7 @@ class LlamaIndexRetrievalOperator(BaseOperator):
embed_model: str | BaseEmbedding | None = None,
llm_conn_id: str | None = None,
embed_conn_id: str | None = None,
+ embedding_kwargs: dict[str, Any] | None = None,
top_k: int = 5,
**kwargs: Any,
) -> None:
@@ -108,6 +116,7 @@ class LlamaIndexRetrievalOperator(BaseOperator):
self.embed_model = embed_model
self.llm_conn_id = llm_conn_id
self.embed_conn_id = embed_conn_id
+ self.embedding_kwargs = embedding_kwargs or {}
self.top_k = top_k
def execute(self, context: Context) -> dict[str, Any]:
@@ -160,6 +169,7 @@ class LlamaIndexRetrievalOperator(BaseOperator):
llm_conn_id=self.llm_conn_id,
embed_conn_id=self.embed_conn_id,
embed_model=self.embed_model,
+ embedding_kwargs=self.embedding_kwargs,
).get_embedding_model()
# ``BaseEmbedding`` always exposes these two methods (see
@@ -169,6 +179,10 @@ class LlamaIndexRetrievalOperator(BaseOperator):
if hasattr(self.embed_model, "get_text_embedding") and hasattr(
self.embed_model, "_get_query_embedding"
):
+ if self.embedding_kwargs:
+ self.log.warning(
+ "embedding_kwargs is ignored when embed_model is a
pre-built embedding model"
+ )
return self.embed_model
raise TypeError(
diff --git a/providers/common/ai/tests/unit/common/ai/hooks/test_langchain.py
b/providers/common/ai/tests/unit/common/ai/hooks/test_langchain.py
index e1af49eef96..58fb66f1bf2 100644
--- a/providers/common/ai/tests/unit/common/ai/hooks/test_langchain.py
+++ b/providers/common/ai/tests/unit/common/ai/hooks/test_langchain.py
@@ -54,6 +54,10 @@ def _conn(password: str = "", host: str = "", extra: dict |
None = None) -> Magi
return mock_conn
+def _init_embeddings(model: str, *, provider: str | None = None, **kwargs):
+ return {"model": model, "provider": provider, "kwargs": kwargs}
+
+
class TestLangChainHookInit:
def test_default_params(self):
hook = LangChainHook()
@@ -61,6 +65,7 @@ class TestLangChainHookInit:
assert hook.embed_conn_id == "langchain_default"
assert hook.llm_model is None
assert hook.embed_model is None
+ assert hook.embedding_kwargs == {}
def test_embed_conn_falls_back_to_llm_conn(self):
hook = LangChainHook(llm_conn_id="my_conn")
@@ -219,6 +224,67 @@ class TestGetEmbeddingModel:
base_url="http://localhost:11434/v1",
)
+ @patch("langchain.embeddings.init_embeddings")
+ @patch.object(LangChainHook, "get_connection")
+ def test_dispatches_with_embedding_kwargs(self, mock_get_conn,
mock_init_embeddings, caplog):
+ mock_get_conn.return_value = _conn(password="sk-test")
+
+ hook = LangChainHook(
+ embed_model="openai:Qwen/Qwen3-Embedding-0.6B",
+ embedding_kwargs={"api_key": "from-kwargs", "dimensions": 128,
"timeout": 30},
+ )
+ hook.get_embedding_model()
+
+ mock_init_embeddings.assert_called_once_with(
+ "openai:Qwen/Qwen3-Embedding-0.6B",
+ api_key="sk-test",
+ dimensions=128,
+ timeout=30,
+ )
+ assert "Connection parameters override embedding_kwargs values:
['api_key']" in caplog.messages
+
+ @pytest.mark.parametrize(
+ ("embed_model", "expect_warning"),
+ [
+ ("openai:text-embedding-3-small", True),
+ ("text-embedding-3-small", False),
+ ],
+ )
+ @patch("langchain.embeddings.init_embeddings")
+ @patch.object(LangChainHook, "get_connection")
+ def test_embedding_kwargs_provider_warning(
+ self, mock_get_conn, mock_init_embeddings, caplog, embed_model,
expect_warning
+ ):
+ mock_get_conn.return_value = _conn()
+ mock_init_embeddings.side_effect = _init_embeddings
+ hook = LangChainHook(
+ embed_model=embed_model,
+ embedding_kwargs={"provider": "custom-provider"},
+ )
+
+ result = hook.get_embedding_model()
+
+ assert result["model"] == embed_model
+ assert result["provider"] == "custom-provider"
+ warning = (
+ "embedding_kwargs['provider'] takes precedence over the provider
prefix in embed_model; "
+ "pass an unprefixed model name"
+ )
+ assert (warning in caplog.messages) is expect_warning
+
+ @patch("langchain.embeddings.init_embeddings")
+ @patch.object(LangChainHook, "get_connection")
+ def test_model_in_embedding_kwargs_raises(self, mock_get_conn,
mock_init_embeddings):
+ mock_get_conn.return_value = _conn()
+ mock_init_embeddings.side_effect = _init_embeddings
+ hook = LangChainHook(
+ embed_model="openai:text-embedding-3-small",
+ embedding_kwargs={"model": "other-model"},
+ )
+
+ with pytest.raises(TypeError, match="multiple values.*model"):
+ hook.get_embedding_model()
+
@patch("langchain.embeddings.init_embeddings")
@patch.object(LangChainHook, "get_connection")
def test_resolves_embed_model_from_extra(self, mock_get_conn,
mock_init_embeddings):
diff --git a/providers/common/ai/tests/unit/common/ai/hooks/test_llamaindex.py
b/providers/common/ai/tests/unit/common/ai/hooks/test_llamaindex.py
index c91866822ce..ac0fcfd71cb 100644
--- a/providers/common/ai/tests/unit/common/ai/hooks/test_llamaindex.py
+++ b/providers/common/ai/tests/unit/common/ai/hooks/test_llamaindex.py
@@ -37,6 +37,7 @@ class TestLlamaIndexHookInit:
assert hook.llm_conn_id == "llamaindex_default"
assert hook.embed_conn_id == "llamaindex_default"
assert hook.embed_model is None
+ assert hook.embedding_kwargs == {}
assert hook.llm_model is None
def test_embed_conn_falls_back_to_llm_conn(self):
@@ -128,6 +129,77 @@ class TestGetEmbeddingModel:
api_base="http://localhost:11434/v1",
)
+ @patch("llama_index.embeddings.openai.OpenAIEmbedding")
+ @patch.object(LlamaIndexHook, "get_connection")
+ def test_dispatches_with_embedding_kwargs(self, mock_get_conn, mock_cls,
caplog):
+ mock_get_conn.return_value = _conn(password="sk-test")
+ mock_cls.model_fields = {"api_key": None, "dimensions": None,
"timeout": None}
+ hook = LlamaIndexHook(
+ embed_model="text-embedding-3-small",
+ embedding_kwargs={"api_key": "from-kwargs", "dimensions": 128,
"timeout": 30},
+ )
+
+ hook.get_embedding_model()
+
+ mock_cls.assert_called_once_with(
+ model="text-embedding-3-small",
+ api_key="sk-test",
+ dimensions=128,
+ timeout=30,
+ )
+ assert "Connection parameters override embedding_kwargs values:
['api_key']" in caplog.messages
+ assert not any(
+ message.startswith("OpenAIEmbedding ignores unsupported
embedding_kwargs")
+ for message in caplog.messages
+ )
+
+ @pytest.mark.parametrize(
+ ("embedding_kwarg", "expect_warning"),
+ [
+ ("http_client", False),
+ ("embeddings_cache", False),
+ ("dimension", True),
+ ("kwargs", True),
+ ],
+ )
+ @patch.object(LlamaIndexHook, "get_connection")
+ def test_warns_about_unsupported_embedding_kwargs(
+ self, mock_get_conn, caplog, embedding_kwarg, expect_warning
+ ):
+ mock_get_conn.return_value = _conn(password="sk-test")
+ hook = LlamaIndexHook(
+ embed_model="text-embedding-3-small",
+ embedding_kwargs={embedding_kwarg: None},
+ )
+
+ hook.get_embedding_model()
+
+ warning = f"OpenAIEmbedding ignores unsupported embedding_kwargs:
['{embedding_kwarg}']"
+ assert (warning in caplog.messages) is expect_warning
+
+ @patch.object(LlamaIndexHook, "get_connection")
+ def test_embedding_kwargs_overrides_model_name(self, mock_get_conn):
+ mock_get_conn.return_value = _conn(password="sk-test")
+ hook = LlamaIndexHook(
+ embed_model="text-embedding-3-small",
+ embedding_kwargs={"model_name": "custom-value"},
+ )
+
+ embedding_model = hook.get_embedding_model()
+
+ assert embedding_model.model_name == "custom-value"
+
+ @patch.object(LlamaIndexHook, "get_connection")
+ def test_model_in_embedding_kwargs_raises(self, mock_get_conn):
+ mock_get_conn.return_value = _conn(password="sk-test")
+ hook = LlamaIndexHook(
+ embed_model="text-embedding-3-small",
+ embedding_kwargs={"model": "other-model"},
+ )
+
+ with pytest.raises(TypeError, match="multiple values.*model"):
+ hook.get_embedding_model()
+
@patch("llama_index.embeddings.openai.OpenAIEmbedding")
@patch.object(LlamaIndexHook, "get_connection")
def test_resolves_model_from_extra(self, mock_get_conn, mock_cls):
diff --git
a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py
b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py
index d3e6b2acfca..e410dfded05 100644
---
a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py
+++
b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_embedding.py
@@ -73,6 +73,7 @@ class TestEmbeddingOperatorInit:
"embed_model",
"llm_conn_id",
"embed_conn_id",
+ "embedding_kwargs",
"persist_dir",
"persist_conn_id",
}
@@ -113,6 +114,7 @@ class TestEmbeddingOperatorExecute:
embed_model="text-embedding-3-small",
llm_conn_id="my_llm_conn",
embed_conn_id="my_embed_conn",
+ embedding_kwargs={"dimensions": 128},
)
op.execute(context=MagicMock())
@@ -120,9 +122,17 @@ class TestEmbeddingOperatorExecute:
llm_conn_id="my_llm_conn",
embed_conn_id="my_embed_conn",
embed_model="text-embedding-3-small",
+ embedding_kwargs={"dimensions": 128},
)
- def test_byo_embed_model_bypasses_hook(self, _li):
+ @pytest.mark.parametrize(
+ ("embedding_kwargs", "expect_warning"),
+ [
+ (None, False),
+ ({"dimensions": 128}, True),
+ ],
+ )
+ def test_byo_embed_model_bypasses_hook(self, _li, caplog,
embedding_kwargs, expect_warning):
# `embed_model` is a non-string instance -> hook is bypassed and the
# user's instance does the embedding.
byo = _byo_embedding(vectors=[[0.5]])
@@ -132,11 +142,14 @@ class TestEmbeddingOperatorExecute:
task_id="test",
documents=[{"text": "doc"}],
embed_model=byo,
+ embedding_kwargs=embedding_kwargs,
)
result = op.execute(context=MagicMock())
byo.get_text_embedding_batch.assert_called_once()
assert result["chunks"][0]["vector"] == [0.5]
+ warning = "embedding_kwargs is ignored when embed_model is a pre-built
embedding model"
+ assert any(warning in record.message for record in caplog.records) is
expect_warning
def test_invalid_embed_model_raises_typeerror(self, _li):
# An object that's neither None/str nor duck-types as BaseEmbedding
diff --git
a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py
b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py
index 58e2bb751c7..c0c4b0c9e89 100644
---
a/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py
+++
b/providers/common/ai/tests/unit/common/ai/operators/test_llamaindex_retrieval.py
@@ -67,6 +67,7 @@ class TestRetrievalOperatorInit:
"embed_model",
"llm_conn_id",
"embed_conn_id",
+ "embedding_kwargs",
}
@@ -135,6 +136,7 @@ class TestRetrievalOperatorOutput:
embed_model="text-embedding-3-small",
llm_conn_id="my_llm_conn",
embed_conn_id="my_embed_conn",
+ embedding_kwargs={"dimensions": 128},
)
op.execute(context=MagicMock())
@@ -142,9 +144,17 @@ class TestRetrievalOperatorOutput:
llm_conn_id="my_llm_conn",
embed_conn_id="my_embed_conn",
embed_model="text-embedding-3-small",
+ embedding_kwargs={"dimensions": 128},
)
- def test_byo_embed_model_bypasses_hook(self, _li, tmp_path):
+ @pytest.mark.parametrize(
+ ("embedding_kwargs", "expect_warning"),
+ [
+ (None, False),
+ ({"dimensions": 128}, True),
+ ],
+ )
+ def test_byo_embed_model_bypasses_hook(self, _li, tmp_path, caplog,
embedding_kwargs, expect_warning):
(tmp_path / "idx").mkdir()
byo = _byo_embedding()
index = _li["load_index_from_storage"].return_value
@@ -155,11 +165,14 @@ class TestRetrievalOperatorOutput:
query="q",
index_persist_dir=str(tmp_path / "idx"),
embed_model=byo,
+ embedding_kwargs=embedding_kwargs,
)
op.execute(context=MagicMock())
kwargs = _li["load_index_from_storage"].call_args.kwargs
assert kwargs["embed_model"] is byo
+ warning = "embedding_kwargs is ignored when embed_model is a pre-built
embedding model"
+ assert any(warning in record.message for record in caplog.records) is
expect_warning
def test_invalid_embed_model_raises_typeerror(self, _li, tmp_path):
# An object that's neither None/str nor duck-types as BaseEmbedding