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 d9a5f778255 Return all embeddings for batched OpenAI inputs (#71989)
d9a5f778255 is described below

commit d9a5f7782554f3be54881bda740ccd2674c61347
Author: Yash jain <[email protected]>
AuthorDate: Thu Sep 17 08:00:53 2026 +0530

    Return all embeddings for batched OpenAI inputs (#71989)
    
    The SDK returns one indexed embedding per batch item, but the hook exposed 
only the first result. Batch callers therefore received a plausible-looking 
response with no indication that every later input had been discarded.
---
 providers/openai/docs/operators/openai.rst         |  3 ++
 .../src/airflow/providers/openai/hooks/openai.py   | 26 +++++++++++++---
 .../airflow/providers/openai/operators/openai.py   |  4 ++-
 .../openai/tests/unit/openai/hooks/test_openai.py  | 35 ++++++++++++++++++++++
 4 files changed, 63 insertions(+), 5 deletions(-)

diff --git a/providers/openai/docs/operators/openai.rst 
b/providers/openai/docs/operators/openai.rst
index 06689ac041d..676b08e45bf 100644
--- a/providers/openai/docs/operators/openai.rst
+++ b/providers/openai/docs/operators/openai.rst
@@ -30,6 +30,9 @@ Using the Operator
 The OpenAIEmbeddingOperator requires the ``input_text`` as an input to 
embedding API. Use the ``conn_id`` parameter to specify the OpenAI connection 
to use to
 connect to your account.
 
+A single string or token array returns one embedding vector. A list of strings 
or token arrays returns
+one vector per input item in the same order.
+
 An example of using the operator:
 
 .. exampleinclude:: /../../openai/tests/system/openai/example_openai.py
diff --git a/providers/openai/src/airflow/providers/openai/hooks/openai.py 
b/providers/openai/src/airflow/providers/openai/hooks/openai.py
index e7af68ca38c..b00fe215015 100644
--- a/providers/openai/src/airflow/providers/openai/hooks/openai.py
+++ b/providers/openai/src/airflow/providers/openai/hooks/openai.py
@@ -20,7 +20,7 @@ from __future__ import annotations
 import time
 from enum import Enum
 from functools import cached_property
-from typing import TYPE_CHECKING, Any, BinaryIO, Literal
+from typing import TYPE_CHECKING, Any, BinaryIO, Literal, overload
 
 from deprecated import deprecated
 from openai import OpenAI
@@ -494,21 +494,39 @@ class OpenAIHook(BaseHook):
         run = self.conn.beta.threads.runs.update(thread_id=thread_id, 
run_id=run_id, **kwargs)
         return run
 
+    @overload
+    def create_embeddings(
+        self,
+        text: str | list[int],
+        model: str = "text-embedding-3-small",
+        **kwargs: Any,
+    ) -> list[float]: ...
+
+    @overload
+    def create_embeddings(
+        self,
+        text: list[str] | list[list[int]],
+        model: str = "text-embedding-3-small",
+        **kwargs: Any,
+    ) -> list[list[float]]: ...
+
     def create_embeddings(
         self,
         text: str | list[str] | list[int] | list[list[int]],
         model: str = "text-embedding-3-small",
         **kwargs: Any,
-    ) -> list[float]:
+    ) -> list[float] | list[list[float]]:
         """
         Generate embeddings for the given text using the given model.
 
         :param text: The text to generate embeddings for.
         :param model: The model to use for generating embeddings.
+        :return: One embedding for a single text or token array; one embedding 
per item for a batch.
         """
         response = self.conn.embeddings.create(model=model, input=text, 
**kwargs)
-        embeddings: list[float] = response.data[0].embedding
-        return embeddings
+        if isinstance(text, str) or (text and isinstance(text[0], int)):
+            return response.data[0].embedding
+        return [item.embedding for item in sorted(response.data, key=lambda 
item: item.index)]
 
     def upload_file(self, file: str, purpose: Literal["fine-tune", 
"assistants", "batch"]) -> FileObject:
         """
diff --git a/providers/openai/src/airflow/providers/openai/operators/openai.py 
b/providers/openai/src/airflow/providers/openai/operators/openai.py
index 2507dd075f2..2c4500358a0 100644
--- a/providers/openai/src/airflow/providers/openai/operators/openai.py
+++ b/providers/openai/src/airflow/providers/openai/operators/openai.py
@@ -40,6 +40,8 @@ class OpenAIEmbeddingOperator(BaseOperator):
     :param model: The OpenAI model to be used for generating the embeddings.
     :param embedding_kwargs: Additional keyword arguments to pass to the 
OpenAI `create_embeddings` method.
 
+    Returns one embedding for a single string or token array, and one 
embedding per item for a batch.
+
     .. seealso::
         For more information on how to use this operator, take a look at the 
guide:
         :ref:`howto/operator:OpenAIEmbeddingOperator`
@@ -68,7 +70,7 @@ class OpenAIEmbeddingOperator(BaseOperator):
         """Return an instance of the OpenAIHook."""
         return OpenAIHook(conn_id=self.conn_id)
 
-    def execute(self, context: Context) -> list[float]:
+    def execute(self, context: Context) -> list[float] | list[list[float]]:
         if not self.input_text or not isinstance(self.input_text, (str, list)):
             raise ValueError(
                 "The 'input_text' must be a non-empty string, list of strings, 
list of integers, or list of lists of integers."
diff --git a/providers/openai/tests/unit/openai/hooks/test_openai.py 
b/providers/openai/tests/unit/openai/hooks/test_openai.py
index 4115f0c2c69..81bcdaa3dca 100644
--- a/providers/openai/tests/unit/openai/hooks/test_openai.py
+++ b/providers/openai/tests/unit/openai/hooks/test_openai.py
@@ -509,6 +509,41 @@ def test_create_embeddings(mock_openai_hook, 
mock_embeddings_response):
     assert embeddings == [0.1, 0.2, 0.3]
 
 
[email protected](
+    ("input_text", "response_items", "expected"),
+    [
+        pytest.param(
+            ["First text", "Second text"],
+            [(1, [0.3, 0.4]), (0, [0.1, 0.2])],
+            [[0.1, 0.2], [0.3, 0.4]],
+            id="text-batch",
+        ),
+        pytest.param(
+            [[1, 2], [3, 4]],
+            [(1, [0.3, 0.4]), (0, [0.1, 0.2])],
+            [[0.1, 0.2], [0.3, 0.4]],
+            id="token-batch",
+        ),
+    ],
+)
+def test_create_batched_embeddings(input_text, response_items, expected):
+    hook = OpenAIHook(conn_id="unused")
+    conn = MagicMock(spec=OpenAI)
+    conn.embeddings.create.return_value = CreateEmbeddingResponse(
+        data=[
+            Embedding(embedding=vector, index=index, object="embedding") for 
index, vector in response_items
+        ],
+        model="text-embedding-3-small",
+        object="list",
+        usage={"prompt_tokens": 4, "total_tokens": 4},
+    )
+    hook.__dict__["conn"] = conn
+
+    embeddings = hook.create_embeddings(input_text)
+
+    assert embeddings == expected
+
+
 @patch("builtins.open", new_callable=mock_open, read_data="test-data")
 def test_upload_file(mock_file_open, mock_openai_hook, mock_file):
     mock_file.name = FILE_NAME

Reply via email to