From 8c99143ccb7269bf354a189ca51c118d05d100ad Mon Sep 17 00:00:00 2001 From: Pankaj Koti Date: Fri, 1 Dec 2023 23:01:36 +0530 Subject: [PATCH 1/3] Bump up openai version to >=1.0 & use get_conn --- airflow/providers/openai/hooks/openai.py | 37 ++++++------- airflow/providers/openai/provider.yaml | 2 +- generated/provider_dependencies.json | 2 +- tests/providers/openai/hooks/test_openai.py | 58 ++++++++++++--------- 4 files changed, 55 insertions(+), 44 deletions(-) diff --git a/airflow/providers/openai/hooks/openai.py b/airflow/providers/openai/hooks/openai.py index fac725b5bea40..3cb961cb83dc3 100644 --- a/airflow/providers/openai/hooks/openai.py +++ b/airflow/providers/openai/hooks/openai.py @@ -17,9 +17,10 @@ from __future__ import annotations +from functools import cached_property from typing import Any -import openai +from openai import OpenAI from airflow.hooks.base import BaseHook @@ -41,13 +42,9 @@ class OpenAIHook(BaseHook): def __init__(self, conn_id: str = default_conn_name, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) self.conn_id = conn_id - openai.api_key = self._get_api_key() - api_base = self._get_api_base() - if api_base: - openai.api_base = api_base - @staticmethod - def get_ui_field_behaviour() -> dict[str, Any]: + @classmethod + def get_ui_field_behaviour(cls) -> dict[str, Any]: """Return custom field behaviour.""" return { "hidden_fields": ["schema", "port", "login", "extra"], @@ -57,21 +54,25 @@ def get_ui_field_behaviour() -> dict[str, Any]: def test_connection(self) -> tuple[bool, str]: try: - openai.Model.list() + self.conn.models.list() return True, "Connection established!" except Exception as e: return False, str(e) - def _get_api_key(self) -> str: - """Get the OpenAI API key from the connection.""" - conn = self.get_connection(self.conn_id) - if not conn.password: - raise ValueError("OpenAI API key not found in connection") - return str(conn.password) + @cached_property + def conn(self) -> OpenAI: + """Return an OpenAI connection object.""" + return self.get_conn() - def _get_api_base(self) -> None | str: + def get_conn(self) -> OpenAI: + """Return an OpenAI connection object.""" conn = self.get_connection(self.conn_id) - return conn.host + url = conn.host or None + password = conn.password + return OpenAI( + api_key=password, + base_url=url, + ) def create_embeddings( self, @@ -84,6 +85,6 @@ def create_embeddings( :param text: The text to generate embeddings for. :param model: The model to use for generating embeddings. """ - response = openai.Embedding.create(model=model, input=text, **kwargs) - embeddings: list[float] = response["data"][0]["embedding"] + response = self.conn.embeddings.create(model=model, input=text, **kwargs) + embeddings: list[float] = response.data[0].embedding return embeddings diff --git a/airflow/providers/openai/provider.yaml b/airflow/providers/openai/provider.yaml index 86226aa3f0f2e..0f9d830a61d6e 100644 --- a/airflow/providers/openai/provider.yaml +++ b/airflow/providers/openai/provider.yaml @@ -39,7 +39,7 @@ integrations: dependencies: - apache-airflow>=2.5.0 - - openai[datalib]>=0.28.1,<1.0 + - openai[datalib]>=1.0 hooks: - integration-name: OpenAI diff --git a/generated/provider_dependencies.json b/generated/provider_dependencies.json index e6405cfe77bdc..e53b5c2054b01 100644 --- a/generated/provider_dependencies.json +++ b/generated/provider_dependencies.json @@ -661,7 +661,7 @@ "openai": { "deps": [ "apache-airflow>=2.5.0", - "openai[datalib]>=0.28.1,<1.0" + "openai[datalib]>=1.0" ], "cross-providers-deps": [], "excluded-python-versions": [] diff --git a/tests/providers/openai/hooks/test_openai.py b/tests/providers/openai/hooks/test_openai.py index cd811107f7988..97aa8776afcc6 100644 --- a/tests/providers/openai/hooks/test_openai.py +++ b/tests/providers/openai/hooks/test_openai.py @@ -16,48 +16,58 @@ # under the License. from __future__ import annotations -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest +from openai import OpenAI +from openai.types import CreateEmbeddingResponse, Embedding from airflow.providers.openai.hooks.openai import OpenAIHook @pytest.fixture def openai_hook(): - with patch("airflow.providers.openai.hooks.openai.OpenAIHook._get_api_key"), patch( - "airflow.providers.openai.hooks.openai.OpenAIHook._get_api_base" - ) as _: + with patch("airflow.providers.openai.hooks.openai.OpenAI") as _: yield OpenAIHook(conn_id="test_conn_id") @pytest.fixture def mock_embeddings_response(): - return {"data": [{"embedding": [0.1, 0.2, 0.3]}]} - - -@pytest.fixture -def mock_completions_response(): - return Mock( - id="completion-id", - object="completion", - created=1234567890, - model="text-davinci-002", - usage={"prompt_tokens": 15, "completion_tokens": 32, "total_tokens": 47}, - choices=[Mock(text="the quick brown fox", finish_reason="stop", index=0)], + return CreateEmbeddingResponse( + data=[Embedding(embedding=[0.1, 0.2, 0.3], index=0, object="embedding")], + model="text-embedding-ada-002-v2", + object="list", + usage={"prompt_tokens": 4, "total_tokens": 4}, ) -def test_create_embeddings(openai_hook, mock_embeddings_response): +@patch("airflow.hooks.base.BaseHook.get_connection") +def test_create_embeddings(mock_get_connection, openai_hook, mock_embeddings_response): text = "Sample text" - with patch("openai.Embedding.create", return_value=mock_embeddings_response): - embeddings = openai_hook.create_embeddings(text) + openai_hook.conn.embeddings.create.return_value = mock_embeddings_response + embeddings = openai_hook.create_embeddings(text) assert embeddings == [0.1, 0.2, 0.3] -def test_get_api_key(): - mock_connection = Mock() - mock_connection.password = "your_api_key" +@patch("openai.OpenAI") +@patch("airflow.hooks.base.BaseHook.get_connection") +def test_openai_hook_get_conn(mock_get_connection, mock_openai): + mock_connection = MagicMock() + mock_connection.host = "http://example.com" + mock_connection.password = "test-api-key" + mock_get_connection.return_value = mock_connection OpenAIHook.get_connection = Mock(return_value=mock_connection) - api_key = OpenAIHook()._get_api_key() - assert api_key == "your_api_key" + + openai_hook = OpenAIHook(conn_id="test_conn") + conn = openai_hook.conn + + assert isinstance(conn, OpenAI) + assert conn.api_key == "test-api-key" + assert conn.base_url == "http://example.com" + + +@patch("openai.OpenAI") +def test_openai_hook_test_connection(mock_openai, openai_hook): + result, message = openai_hook.test_connection() + assert result is True + assert message == "Connection established!" From d7524ba06fe0dea3a5a7c3f4825d47fcb48067a3 Mon Sep 17 00:00:00 2001 From: Pankaj Koti Date: Sun, 3 Dec 2023 13:09:55 +0530 Subject: [PATCH 2/3] Remove 'as _' for with patch as it's not needed --- tests/providers/openai/hooks/test_openai.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/providers/openai/hooks/test_openai.py b/tests/providers/openai/hooks/test_openai.py index 97aa8776afcc6..e0be3e4bee21a 100644 --- a/tests/providers/openai/hooks/test_openai.py +++ b/tests/providers/openai/hooks/test_openai.py @@ -27,7 +27,7 @@ @pytest.fixture def openai_hook(): - with patch("airflow.providers.openai.hooks.openai.OpenAI") as _: + with patch("airflow.providers.openai.hooks.openai.OpenAI"): yield OpenAIHook(conn_id="test_conn_id") From f261609a96888b9e0ee600683ad37000f19fdeda Mon Sep 17 00:00:00 2001 From: Pankaj Koti Date: Tue, 5 Dec 2023 18:24:28 +0530 Subject: [PATCH 3/3] Address @ephraimbuddy's comment --- airflow/providers/openai/hooks/openai.py | 13 +- .../connections.rst | 17 +++ tests/providers/openai/hooks/test_openai.py | 119 ++++++++++++++---- 3 files changed, 118 insertions(+), 31 deletions(-) diff --git a/airflow/providers/openai/hooks/openai.py b/airflow/providers/openai/hooks/openai.py index 3cb961cb83dc3..f57c41c2b9f4b 100644 --- a/airflow/providers/openai/hooks/openai.py +++ b/airflow/providers/openai/hooks/openai.py @@ -47,7 +47,7 @@ def __init__(self, conn_id: str = default_conn_name, *args: Any, **kwargs: Any) def get_ui_field_behaviour(cls) -> dict[str, Any]: """Return custom field behaviour.""" return { - "hidden_fields": ["schema", "port", "login", "extra"], + "hidden_fields": ["schema", "port", "login"], "relabeling": {"password": "API Key"}, "placeholders": {}, } @@ -67,11 +67,14 @@ def conn(self) -> OpenAI: def get_conn(self) -> OpenAI: """Return an OpenAI connection object.""" conn = self.get_connection(self.conn_id) - url = conn.host or None - password = conn.password + extras = conn.extra_dejson + openai_client_kwargs = extras.get("openai_client_kwargs", {}) + api_key = openai_client_kwargs.pop("api_key", None) or conn.password + base_url = openai_client_kwargs.pop("base_url", None) or conn.host or None return OpenAI( - api_key=password, - base_url=url, + api_key=api_key, + base_url=base_url, + **openai_client_kwargs, ) def create_embeddings( diff --git a/docs/apache-airflow-providers-openai/connections.rst b/docs/apache-airflow-providers-openai/connections.rst index 88e79df59d6d3..8ef7ee456b1c9 100644 --- a/docs/apache-airflow-providers-openai/connections.rst +++ b/docs/apache-airflow-providers-openai/connections.rst @@ -35,3 +35,20 @@ API Key (required) Host (optional) The host address of the OpenAI instance. + +Extra (optional) + Specify the extra parameters (as json dictionary) that can be used in the + connection. All parameters are optional. + This ``extra`` field accepts a nested dictionary with key ``openai_client_kwargs`` as key-value pairs that + are passed to the `OpenAI client `__ + on instantiation. For example, to set the timeout for the client, you can pass the following dictionary + as the ``extra`` field: + + .. code-block:: json + + { + "openai_client_kwargs": { + "timeout": 10, + "api_key": "YOUR_API_KEY" + } + } diff --git a/tests/providers/openai/hooks/test_openai.py b/tests/providers/openai/hooks/test_openai.py index e0be3e4bee21a..a80be35dfbb9d 100644 --- a/tests/providers/openai/hooks/test_openai.py +++ b/tests/providers/openai/hooks/test_openai.py @@ -16,19 +16,31 @@ # under the License. from __future__ import annotations -from unittest.mock import MagicMock, Mock, patch +import os +from unittest.mock import patch import pytest -from openai import OpenAI from openai.types import CreateEmbeddingResponse, Embedding +from airflow.models import Connection from airflow.providers.openai.hooks.openai import OpenAIHook @pytest.fixture -def openai_hook(): +def mock_openai_connection(): + conn_id = "openai_conn" + conn = Connection( + conn_id=conn_id, + conn_type="openai", + ) + os.environ[f"AIRFLOW_CONN_{conn.conn_id.upper()}"] = conn.get_uri() + yield conn + + +@pytest.fixture +def mock_openai_hook(mock_openai_connection): with patch("airflow.providers.openai.hooks.openai.OpenAI"): - yield OpenAIHook(conn_id="test_conn_id") + yield OpenAIHook(conn_id=mock_openai_connection.conn_id) @pytest.fixture @@ -41,33 +53,88 @@ def mock_embeddings_response(): ) -@patch("airflow.hooks.base.BaseHook.get_connection") -def test_create_embeddings(mock_get_connection, openai_hook, mock_embeddings_response): +def test_create_embeddings(mock_openai_hook, mock_embeddings_response): text = "Sample text" - openai_hook.conn.embeddings.create.return_value = mock_embeddings_response - embeddings = openai_hook.create_embeddings(text) + mock_openai_hook.conn.embeddings.create.return_value = mock_embeddings_response + embeddings = mock_openai_hook.create_embeddings(text) assert embeddings == [0.1, 0.2, 0.3] -@patch("openai.OpenAI") -@patch("airflow.hooks.base.BaseHook.get_connection") -def test_openai_hook_get_conn(mock_get_connection, mock_openai): - mock_connection = MagicMock() - mock_connection.host = "http://example.com" - mock_connection.password = "test-api-key" - mock_get_connection.return_value = mock_connection - OpenAIHook.get_connection = Mock(return_value=mock_connection) +def test_openai_hook_test_connection(mock_openai_hook): + result, message = mock_openai_hook.test_connection() + assert result is True + assert message == "Connection established!" - openai_hook = OpenAIHook(conn_id="test_conn") - conn = openai_hook.conn - assert isinstance(conn, OpenAI) - assert conn.api_key == "test-api-key" - assert conn.base_url == "http://example.com" +@patch("airflow.providers.openai.hooks.openai.OpenAI") +def test_get_conn_with_api_key_in_extra(mock_client): + conn_id = "api_key_in_extra" + conn = Connection( + conn_id=conn_id, + conn_type="openai", + extra={"openai_client_kwargs": {"api_key": "api_key_in_extra"}}, + ) + os.environ[f"AIRFLOW_CONN_{conn.conn_id.upper()}"] = conn.get_uri() + hook = OpenAIHook(conn_id=conn_id) + hook.get_conn() + mock_client.assert_called_once_with( + api_key="api_key_in_extra", + base_url=None, + ) -@patch("openai.OpenAI") -def test_openai_hook_test_connection(mock_openai, openai_hook): - result, message = openai_hook.test_connection() - assert result is True - assert message == "Connection established!" +@patch("airflow.providers.openai.hooks.openai.OpenAI") +def test_get_conn_with_api_key_in_password(mock_client): + conn_id = "api_key_in_password" + conn = Connection( + conn_id=conn_id, + conn_type="openai", + password="api_key_in_password", + ) + os.environ[f"AIRFLOW_CONN_{conn.conn_id.upper()}"] = conn.get_uri() + hook = OpenAIHook(conn_id=conn_id) + hook.get_conn() + mock_client.assert_called_once_with( + api_key="api_key_in_password", + base_url=None, + ) + + +@patch("airflow.providers.openai.hooks.openai.OpenAI") +def test_get_conn_with_base_url_in_extra(mock_client): + conn_id = "base_url_in_extra" + conn = Connection( + conn_id=conn_id, + conn_type="openai", + extra={"openai_client_kwargs": {"base_url": "base_url_in_extra", "api_key": "api_key_in_extra"}}, + ) + os.environ[f"AIRFLOW_CONN_{conn.conn_id.upper()}"] = conn.get_uri() + hook = OpenAIHook(conn_id=conn_id) + hook.get_conn() + mock_client.assert_called_once_with( + api_key="api_key_in_extra", + base_url="base_url_in_extra", + ) + + +@patch("airflow.providers.openai.hooks.openai.OpenAI") +def test_get_conn_with_openai_client_kwargs(mock_client): + conn_id = "openai_client_kwargs" + conn = Connection( + conn_id=conn_id, + conn_type="openai", + extra={ + "openai_client_kwargs": { + "api_key": "api_key_in_extra", + "organization": "organization_in_extra", + } + }, + ) + os.environ[f"AIRFLOW_CONN_{conn.conn_id.upper()}"] = conn.get_uri() + hook = OpenAIHook(conn_id=conn_id) + hook.get_conn() + mock_client.assert_called_once_with( + api_key="api_key_in_extra", + base_url=None, + organization="organization_in_extra", + )