diff --git a/litellm/llms/databricks/common_utils.py b/litellm/llms/databricks/common_utils.py index 608f29a03a7..afbf7c5df81 100644 --- a/litellm/llms/databricks/common_utils.py +++ b/litellm/llms/databricks/common_utils.py @@ -387,8 +387,10 @@ class DatabricksBase: if api_key is not None: headers["Authorization"] = f"Bearer {api_key}" - # Set User-Agent with optional custom prefix - headers["User-Agent"] = self._build_user_agent(custom_user_agent) + # Set User-Agent with optional custom prefix, but don't overwrite + # if the caller explicitly set one (e.g. via extra_headers) + if "User-Agent" not in headers: + headers["User-Agent"] = self._build_user_agent(custom_user_agent) # Debug logging with redaction (never log actual tokens) verbose_logger.debug( diff --git a/tests/test_litellm/llms/databricks/test_databricks_partner_integration.py b/tests/test_litellm/llms/databricks/test_databricks_partner_integration.py index 800066ac5bf..aa7901598dd 100644 --- a/tests/test_litellm/llms/databricks/test_databricks_partner_integration.py +++ b/tests/test_litellm/llms/databricks/test_databricks_partner_integration.py @@ -334,6 +334,40 @@ class TestValidateEnvironmentUserAgent: assert headers["User-Agent"].startswith("mycompany_litellm/") + def test_extra_headers_user_agent_preserved(self, monkeypatch): + """User-Agent set via headers dict (extra_headers) is not overwritten.""" + monkeypatch.delenv("DATABRICKS_CLIENT_ID", raising=False) + monkeypatch.delenv("DATABRICKS_CLIENT_SECRET", raising=False) + + databricks_base = DatabricksBase() + + api_base, headers = databricks_base.databricks_validate_environment( + api_key="test-key", + api_base="https://adb-123.net/serving-endpoints", + endpoint_type="chat_completions", + custom_endpoint=False, + headers={"User-Agent": "My-Custom-Agent/1.0"}, + ) + + assert headers["User-Agent"] == "My-Custom-Agent/1.0" + + def test_default_user_agent_when_not_in_headers(self, monkeypatch): + """Default User-Agent is set when headers dict has no User-Agent.""" + monkeypatch.delenv("DATABRICKS_CLIENT_ID", raising=False) + monkeypatch.delenv("DATABRICKS_CLIENT_SECRET", raising=False) + + databricks_base = DatabricksBase() + + api_base, headers = databricks_base.databricks_validate_environment( + api_key="test-key", + api_base="https://adb-123.net/serving-endpoints", + endpoint_type="chat_completions", + custom_endpoint=False, + headers={}, + ) + + assert headers["User-Agent"].startswith("litellm/") + class TestSDKPartnerTelemetry: """Test that SDK partner telemetry is registered."""