From 4fa28d2fda67637aceb70bb7f64df83e1b4939c4 Mon Sep 17 00:00:00 2001 From: Bill Wang Date: Fri, 25 Sep 2026 16:13:09 +0000 Subject: [PATCH] Add model_services OpenAI query helper Adds a ModelServicesExt mixin so w.model_services.get_open_ai_client() queries Unity Catalog v3 system.ai.* model services via the AI Gateway, mirroring the existing ServingEndpointsExt query surface. Co-authored-by: Isaac Signed-off-by: Bill Wang --- databricks/sdk/__init__.py | 7 ++ databricks/sdk/mixins/model_services.py | 107 ++++++++++++++++++++++++ tests/test_model_services_mixin.py | 77 +++++++++++++++++ 3 files changed, 191 insertions(+) create mode 100644 databricks/sdk/mixins/model_services.py create mode 100644 tests/test_model_services_mixin.py diff --git a/databricks/sdk/__init__.py b/databricks/sdk/__init__.py index 5e670ff81..41f6e664d 100644 --- a/databricks/sdk/__init__.py +++ b/databricks/sdk/__init__.py @@ -54,6 +54,7 @@ from databricks.sdk.mixins.compute import ClustersExt from databricks.sdk.mixins.workspace import WorkspaceExt from databricks.sdk.mixins.open_ai_client import ServingEndpointsExt +from databricks.sdk.mixins.model_services import ModelServicesExt from databricks.sdk.mixins.jobs import JobsExt from databricks.sdk.oauth import AuthorizationDetail from databricks.sdk.service.iam import AccessControlAPI @@ -432,6 +433,7 @@ def __init__( self._materialized_features = pkg_ml.MaterializedFeaturesAPI(self._api_client) self._metastores = pkg_catalog.MetastoresAPI(self._api_client) self._model_registry = pkg_ml.ModelRegistryAPI(self._api_client) + self._model_services = ModelServicesExt(self._api_client) self._model_versions = pkg_catalog.ModelVersionsAPI(self._api_client) self._notification_destinations = pkg_settings.NotificationDestinationsAPI(self._api_client) self._online_tables = pkg_catalog.OnlineTablesAPI(self._api_client) @@ -860,6 +862,11 @@ def model_registry(self) -> pkg_ml.ModelRegistryAPI: """Note: This API reference documents APIs for the Workspace Model Registry.""" return self._model_registry + @property + def model_services(self) -> ModelServicesExt: + """Model services (model APIs) are Unity Catalog securables representing governed LLM endpoints, queried through Unity Gateway.""" + return self._model_services + @property def model_versions(self) -> pkg_catalog.ModelVersionsAPI: """Databricks provides a hosted version of MLflow Model Registry in Unity Catalog.""" diff --git a/databricks/sdk/mixins/model_services.py b/databricks/sdk/mixins/model_services.py new file mode 100644 index 000000000..91a68439f --- /dev/null +++ b/databricks/sdk/mixins/model_services.py @@ -0,0 +1,107 @@ +class ModelServicesExt: + """Extension for querying model services (model APIs) through Unity Gateway. + + A model service is a Unity Catalog securable that represents a governed LLM endpoint. + """ + + def __init__(self, api_client): + self._api = api_client + + # Using the HTTP Client to pass in the databricks authorization + # This method will be called on every invocation, so when using with model serving will always get the refreshed token + def _get_authorized_http_client(self): + import httpx + + class BearerAuth(httpx.Auth): + + def __init__(self, get_headers_func): + self.get_headers_func = get_headers_func + + def auth_flow(self, request: httpx.Request) -> httpx.Request: + auth_headers = self.get_headers_func() + request.headers["Authorization"] = auth_headers["Authorization"] + yield request + + databricks_token_auth = BearerAuth(self._api._cfg.authenticate) + + # Create an HTTP client with Bearer Token authentication + http_client = httpx.Client(auth=databricks_token_auth) + return http_client + + def get_open_ai_client(self, **kwargs): + """Create an OpenAI client configured for querying model services through Unity Gateway. + + Returns an OpenAI client instance that is pre-configured to send requests to Unity + Gateway. The client uses Databricks authentication to query model services within the + workspace associated with the current WorkspaceClient instance. A model service is + queried by passing its Unity Catalog three-level name as the ``model`` parameter. + + Args: + **kwargs: Additional parameters to pass to the OpenAI client constructor. + Common parameters include: + - timeout (float): Request timeout in seconds (e.g., 30.0) + - max_retries (int): Maximum number of retries for failed requests (e.g., 3) + - default_headers (dict): Additional headers to include with requests + - default_query (dict): Additional query parameters to include with requests + + Any parameter accepted by the OpenAI client constructor can be passed here, + except for the following parameters which are reserved for Databricks integration: + base_url, api_key, http_client + + Returns: + OpenAI: An OpenAI client instance configured for querying model services. + + Raises: + ImportError: If the OpenAI library is not installed. + ValueError: If any reserved Databricks parameters are provided in kwargs. + + Example: + >>> client = workspace_client.model_services.get_open_ai_client() + >>> client.chat.completions.create( + ... model="system.ai.claude-sonnet-4-5", + ... messages=[{"role": "user", "content": "Hello!"}], + ... max_tokens=256, + ... ) + """ + try: + from openai import OpenAI + except Exception: + raise ImportError( + "Open AI is not installed. Please install the Databricks SDK with the following command `pip install databricks-sdk[openai]`" + ) + + # Check for reserved parameters that should not be overridden + reserved_params = {"base_url", "api_key", "http_client"} + conflicting_params = reserved_params.intersection(kwargs.keys()) + if conflicting_params: + raise ValueError( + f"Cannot override reserved Databricks parameters: {', '.join(sorted(conflicting_params))}. " + f"These parameters are automatically configured for Databricks Model Serving." + ) + + # Default parameters that are required for Databricks integration + client_params = { + "base_url": self._api._cfg.host + "/ai-gateway/mlflow/v1", + "api_key": "no-token", # Passing in a placeholder to pass validations, this will not be used + "http_client": self._get_authorized_http_client(), + } + + # Update with any additional parameters passed by the user + client_params.update(kwargs) + + return OpenAI(**client_params) + + def get_langchain_chat_open_ai_client(self, model): + try: + from langchain_openai import ChatOpenAI + except Exception: + raise ImportError( + "Langchain Open AI is not installed. Please install the Databricks SDK with the following command `pip install databricks-sdk[openai]` and ensure you are using python>3.7" + ) + + return ChatOpenAI( + model=model, + openai_api_base=self._api._cfg.host + "/ai-gateway/mlflow/v1", + api_key="no-token", # Passing in a placeholder to pass validations, this will not be used + http_client=self._get_authorized_http_client(), + ) diff --git a/tests/test_model_services_mixin.py b/tests/test_model_services_mixin.py new file mode 100644 index 000000000..576645f67 --- /dev/null +++ b/tests/test_model_services_mixin.py @@ -0,0 +1,77 @@ +import sys + +import pytest + +from databricks.sdk.core import Config + + +def test_open_ai_client(monkeypatch): + from databricks.sdk import WorkspaceClient + + monkeypatch.setenv("DATABRICKS_HOST", "test_host") + monkeypatch.setenv("DATABRICKS_TOKEN", "test_token") + w = WorkspaceClient(config=Config()) + client = w.model_services.get_open_ai_client() + + assert client.base_url == "https://test_host/ai-gateway/mlflow/v1/" + assert client.api_key == "no-token" + + +def test_open_ai_client_with_custom_params(monkeypatch): + from databricks.sdk import WorkspaceClient + + monkeypatch.setenv("DATABRICKS_HOST", "test_host") + monkeypatch.setenv("DATABRICKS_TOKEN", "test_token") + w = WorkspaceClient(config=Config()) + + client = w.model_services.get_open_ai_client(timeout=30.0, max_retries=3) + + assert client.base_url == "https://test_host/ai-gateway/mlflow/v1/" + assert client.api_key == "no-token" + assert client.timeout == 30.0 + assert client.max_retries == 3 + + +def test_open_ai_client_prevents_reserved_param_override(monkeypatch): + from databricks.sdk import WorkspaceClient + + monkeypatch.setenv("DATABRICKS_HOST", "test_host") + monkeypatch.setenv("DATABRICKS_TOKEN", "test_token") + w = WorkspaceClient(config=Config()) + + with pytest.raises(ValueError, match="Cannot override reserved Databricks parameters: base_url"): + w.model_services.get_open_ai_client(base_url="https://custom-host") + + with pytest.raises(ValueError, match="Cannot override reserved Databricks parameters: api_key"): + w.model_services.get_open_ai_client(api_key="custom-key") + + with pytest.raises(ValueError, match="Cannot override reserved Databricks parameters: http_client"): + w.model_services.get_open_ai_client(http_client=None) + + with pytest.raises(ValueError, match="Cannot override reserved Databricks parameters: api_key, base_url"): + w.model_services.get_open_ai_client(base_url="https://custom-host", api_key="custom-key") + + +def test_open_ai_client_uses_config_authenticate(monkeypatch): + from databricks.sdk import WorkspaceClient + + monkeypatch.setenv("DATABRICKS_HOST", "test_host") + monkeypatch.setenv("DATABRICKS_TOKEN", "test_token") + w = WorkspaceClient(config=Config()) + client = w.model_services.get_open_ai_client() + + auth_headers = client._client._auth.get_headers_func() + assert auth_headers["Authorization"] == "Bearer test_token" + + +@pytest.mark.skipif(sys.version_info < (3, 8), reason="Requires Python > 3.7") +def test_langchain_open_ai_client(monkeypatch): + from databricks.sdk import WorkspaceClient + + monkeypatch.setenv("DATABRICKS_HOST", "test_host") + monkeypatch.setenv("DATABRICKS_TOKEN", "test_token") + w = WorkspaceClient(config=Config()) + client = w.model_services.get_langchain_chat_open_ai_client("system.ai.claude-sonnet-4-5") + + assert client.openai_api_base == "https://test_host/ai-gateway/mlflow/v1" + assert client.model_name == "system.ai.claude-sonnet-4-5"