diff --git a/integrations/openai/src/databricks_openai/__init__.py b/integrations/openai/src/databricks_openai/__init__.py index ce04d0c0..34948429 100644 --- a/integrations/openai/src/databricks_openai/__init__.py +++ b/integrations/openai/src/databricks_openai/__init__.py @@ -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 `_ 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__ = [ @@ -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__)) diff --git a/integrations/openai/tests/unit_tests/test_imports.py b/integrations/openai/tests/unit_tests/test_imports.py index 9151fe77..39188ae0 100644 --- a/integrations/openai/tests/unit_tests/test_imports.py +++ b/integrations/openai/tests/unit_tests/test_imports.py @@ -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 diff --git a/integrations/openai/tests/unit_tests/test_lazy_imports.py b/integrations/openai/tests/unit_tests/test_lazy_imports.py new file mode 100644 index 00000000..8a9d6182 --- /dev/null +++ b/integrations/openai/tests/unit_tests/test_lazy_imports.py @@ -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