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

Reply via email to