fix(databricks): route Unity model services through AI Gateway (#40492)

Co-authored-by: Claude Code <noreply@anthropic.com>
This commit is contained in:
tin-berri 2026-09-09 17:15:48 -07:00 • committed by GitHub
parent a650178ebe
commit 6b721de3e5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 101 additions and 9 deletions

View file

@ -250,8 +250,10 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
litellm_params: dict,
stream: bool | None = None,
) -> str:
api_base = self._get_api_base(api_base)
complete_url: Final = f"{api_base}/chat/completions"
use_ai_gateway: Final = model.removeprefix("databricks/").count(".") >= 2
api_base = self._get_api_base(api_base, use_ai_gateway=use_ai_gateway)
url_base: Final = api_base.rstrip("/") if use_ai_gateway else api_base
complete_url: Final = f"{url_base}/chat/completions"
return complete_url
def get_supported_openai_params(self, model: str | None = None) -> list:

View file

@ -177,19 +177,13 @@ class DatabricksBase:
# Default: just litellm
return f"litellm/{version}"
def _get_api_base(self, api_base: str | None) -> str:
"""
Get the Databricks API base URL.
If not provided, attempts to get it from the Databricks SDK.
"""
def _get_api_base(self, api_base: str | None, use_ai_gateway: bool = False) -> str:
if api_base is None:
try:
from databricks.sdk import WorkspaceClient
databricks_client: Final = WorkspaceClient()
api_base = f"{databricks_client.config.host}/serving-endpoints"
return api_base
except ImportError:
raise DatabricksException(
status_code=400,
@ -198,6 +192,18 @@ class DatabricksBase:
"or install the databricks-sdk Python library."
),
)
if not use_ai_gateway:
return api_base
normalized_api_base: Final = api_base.rstrip("/")
if normalized_api_base.endswith("/ai-gateway/mlflow/v1"):
return normalized_api_base
if normalized_api_base.endswith("/serving-endpoints"):
return f"{normalized_api_base.removesuffix('/serving-endpoints')}/ai-gateway/mlflow/v1"
api_base_parts: Final = urlsplit(normalized_api_base)
if api_base_parts.path in ("", "/"):
return f"{normalized_api_base}/ai-gateway/mlflow/v1"
return api_base
def _get_oauth_m2m_token(

View file

@ -255,6 +255,19 @@ def test_transform_messages_sanitizes_empty_content():
assert result[1]["content"] == "Hi"
def test_transform_request_preserves_unity_model_service_name():
config = DatabricksConfig()
result = config.transform_request(
model="system.ai.kimi-k3",
messages=[{"role": "user", "content": "hello"}],
optional_params={},
litellm_params={},
headers={},
)
assert result["model"] == "system.ai.kimi-k3"
def test_transform_request_strips_thinking_blocks_and_reasoning_content():
"""Regression for LIT-6762: replaying an assistant turn that litellm decorated with
`thinking_blocks` / `reasoning_content` made Databricks 400 with

View file

@ -657,6 +657,77 @@ class TestEndpointURLConstruction:
assert api_base.endswith("/chat/completions")
def test_chat_gateway_endpoint_for_unity_model_on_legacy_base(self, monkeypatch):
from litellm.llms.databricks.chat.transformation import DatabricksConfig
monkeypatch.delenv("DATABRICKS_CLIENT_ID", raising=False)
monkeypatch.delenv("DATABRICKS_CLIENT_SECRET", raising=False)
url = DatabricksConfig().get_complete_url(
api_base="https://test.net/serving-endpoints",
api_key="test-key",
model="system.ai.kimi-k3",
optional_params={},
litellm_params={},
)
assert url == "https://test.net/ai-gateway/mlflow/v1/chat/completions"
def test_chat_gateway_endpoint_preserves_explicit_gateway_base(self, monkeypatch):
from litellm.llms.databricks.chat.transformation import DatabricksConfig
monkeypatch.delenv("DATABRICKS_CLIENT_ID", raising=False)
monkeypatch.delenv("DATABRICKS_CLIENT_SECRET", raising=False)
url = DatabricksConfig().get_complete_url(
api_base="https://test.net/ai-gateway/mlflow/v1/",
api_key="test-key",
model="system.ai.kimi-k3",
optional_params={},
litellm_params={},
)
assert url == "https://test.net/ai-gateway/mlflow/v1/chat/completions"
def test_chat_gateway_preserves_unity_model_service_name_with_explicit_base(self, monkeypatch):
from litellm.llms.databricks.chat.transformation import DatabricksConfig
monkeypatch.delenv("DATABRICKS_CLIENT_ID", raising=False)
monkeypatch.delenv("DATABRICKS_CLIENT_SECRET", raising=False)
config = DatabricksConfig()
request = config.transform_request(
model="catalog.schema.kimi-k3",
messages=[{"role": "user", "content": "hello"}],
optional_params={},
litellm_params={},
headers={},
)
assert config.get_complete_url(
api_base="https://test.net/ai-gateway/mlflow/v1",
api_key="test-key",
model="catalog.schema.kimi-k3",
optional_params={},
litellm_params={},
) == "https://test.net/ai-gateway/mlflow/v1/chat/completions"
assert request["model"] == "catalog.schema.kimi-k3"
def test_chat_legacy_endpoint_remains_default(self, monkeypatch):
from litellm.llms.databricks.chat.transformation import DatabricksConfig
monkeypatch.delenv("DATABRICKS_CLIENT_ID", raising=False)
monkeypatch.delenv("DATABRICKS_CLIENT_SECRET", raising=False)
url = DatabricksConfig().get_complete_url(
api_base="https://test.net/serving-endpoints",
api_key="test-key",
model="databricks-kimi-k3",
optional_params={},
litellm_params={},
)
assert url == "https://test.net/serving-endpoints/chat/completions"
def test_embeddings_endpoint(self, monkeypatch):
"""Embeddings endpoint is correctly appended."""
monkeypatch.delenv("DATABRICKS_CLIENT_ID", raising=False)