mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(databricks): route Unity model services through AI Gateway (#40492)
Co-authored-by: Claude Code <noreply@anthropic.com>
This commit is contained in:
parent
a650178ebe
commit
6b721de3e5
4 changed files with 101 additions and 9 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue