Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 39 additions & 7 deletions integrations/openai/src/databricks_openai/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,17 +8,17 @@
- :class:`databricks_openai.UCFunctionToolkit`
- :class:`databricks_openai.DatabricksFunctionClient`
- :func:`databricks_openai.set_uc_function_client`
- :class:`databricks_openai.DatabricksOpenAI`
- :class:`databricks_openai.AsyncDatabricksOpenAI`
- :class:`databricks_openai.McpServerToolkit`
- :class:`databricks_openai.ToolInfo`
- :class:`databricks_openai.VectorSearchRetrieverTool`

Refer to the Unity Catalog `documentation <https://docs.unitycatalog.io/ai/integrations/openai/#using-unity-catalog-ai-with-the-openai-sdk>`_ for more information.
"""

from unitycatalog.ai.core.base import set_uc_function_client
from unitycatalog.ai.core.databricks import DatabricksFunctionClient
from unitycatalog.ai.openai.toolkit import UCFunctionToolkit

from databricks_openai.mcp_server_toolkit import McpServerToolkit, ToolInfo
from databricks_openai.utils.clients import AsyncDatabricksOpenAI, DatabricksOpenAI
from databricks_openai.vector_search_retriever_tool import VectorSearchRetrieverTool
from importlib import import_module
from typing import Any

# Expose all integrations to users under databricks-openai
__all__ = [
Expand All @@ -31,3 +31,35 @@
"McpServerToolkit",
"ToolInfo",
]

_LAZY_EXPORTS = {
"DatabricksOpenAI": ("databricks_openai.utils.clients", "DatabricksOpenAI"),
"AsyncDatabricksOpenAI": ("databricks_openai.utils.clients", "AsyncDatabricksOpenAI"),
"McpServerToolkit": ("databricks_openai.mcp_server_toolkit", "McpServerToolkit"),
"ToolInfo": ("databricks_openai.mcp_server_toolkit", "ToolInfo"),
"VectorSearchRetrieverTool": (
"databricks_openai.vector_search_retriever_tool",
"VectorSearchRetrieverTool",
),
"UCFunctionToolkit": ("unitycatalog.ai.openai.toolkit", "UCFunctionToolkit"),
"DatabricksFunctionClient": (
"unitycatalog.ai.core.databricks",
"DatabricksFunctionClient",
),
"set_uc_function_client": ("unitycatalog.ai.core.base", "set_uc_function_client"),
}


def __getattr__(name: str) -> Any:
try:
module_name, attribute_name = _LAZY_EXPORTS[name]
except KeyError as exc:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from exc

value = getattr(import_module(module_name), attribute_name)
globals()[name] = value
return value


def __dir__() -> list[str]:
return sorted(set(globals()) | set(__all__))
41 changes: 21 additions & 20 deletions integrations/openai/tests/unit_tests/test_imports.py
Original file line number Diff line number Diff line change
@@ -1,21 +1,22 @@
from databricks_openai import (
AsyncDatabricksOpenAI,
DatabricksFunctionClient,
DatabricksOpenAI,
McpServerToolkit,
ToolInfo,
UCFunctionToolkit,
VectorSearchRetrieverTool,
set_uc_function_client,
)
from databricks_openai.agents import McpServer
def test_package_root_reexports_remain_available():
from databricks_openai import (
AsyncDatabricksOpenAI,
DatabricksFunctionClient,
DatabricksOpenAI,
McpServerToolkit,
ToolInfo,
UCFunctionToolkit,
VectorSearchRetrieverTool,
set_uc_function_client,
)
from databricks_openai.agents import McpServer

assert DatabricksFunctionClient
assert UCFunctionToolkit
assert VectorSearchRetrieverTool
assert set_uc_function_client
assert DatabricksOpenAI
assert AsyncDatabricksOpenAI
assert McpServerToolkit
assert ToolInfo
assert McpServer
assert DatabricksFunctionClient
assert UCFunctionToolkit
assert VectorSearchRetrieverTool
assert set_uc_function_client
assert DatabricksOpenAI
assert AsyncDatabricksOpenAI
assert McpServerToolkit
assert ToolInfo
assert McpServer
37 changes: 37 additions & 0 deletions integrations/openai/tests/unit_tests/test_lazy_imports.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
import importlib
import sys

FORBIDDEN_EAGER_IMPORTS = [
"databricks_openai.mcp_server_toolkit",
"databricks_openai.vector_search_retriever_tool",
"unitycatalog.ai.core.base",
"unitycatalog.ai.core.databricks",
"unitycatalog.ai.openai.toolkit",
]


def test_package_root_import_does_not_eagerly_import_integrations(monkeypatch):
monkeypatch.delitem(sys.modules, "databricks_openai", raising=False)
for module_name in FORBIDDEN_EAGER_IMPORTS:
monkeypatch.delitem(sys.modules, module_name, raising=False)

importlib.import_module("databricks_openai")

loaded = [name for name in FORBIDDEN_EAGER_IMPORTS if name in sys.modules]
assert not loaded


def test_package_root_client_import_only_loads_client_module(monkeypatch):
monkeypatch.delitem(sys.modules, "databricks_openai", raising=False)
monkeypatch.delitem(sys.modules, "databricks_openai.utils.clients", raising=False)
for module_name in FORBIDDEN_EAGER_IMPORTS:
monkeypatch.delitem(sys.modules, module_name, raising=False)

databricks_openai = importlib.import_module("databricks_openai")
client = databricks_openai.DatabricksOpenAI

assert client.__name__ == "DatabricksOpenAI"
assert "databricks_openai.utils.clients" in sys.modules

loaded = [name for name in FORBIDDEN_EAGER_IMPORTS if name in sys.modules]
assert not loaded