Skip to content

Commit 8cf89b6

Browse files
Support OAuth M2M (client credentials) on the default backend; reject unknown auth_type
oauth_client_id + oauth_client_secret now authenticate a Databricks service principal with the client-credentials flow against the workspace /oidc/v1/token endpoint (scope all-apis), with tokens refreshed as they expire. Previously the Thrift backend ignored oauth_client_secret and started an interactive browser login with the service principal's client ID, which blocks indefinitely on a server. This matches the credential shape the kernel backend already accepts. An auth_type the connector does not implement now raises ValueError instead of falling through to the interactive login, and a U2M auth_type combined with oauth_client_secret is rejected as ambiguous. The client-credentials token source also ignores extra fields in the token response (the workspace endpoint returns 'scope'), which made OAuthResponse(**payload) raise.
1 parent 01564c7 commit 8cf89b6

6 files changed

Lines changed: 153 additions & 1 deletion

File tree

‎src/databricks/sql/auth/auth.py‎

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
ExternalAuthProvider,
77
DatabricksOAuthProvider,
88
AzureServicePrincipalCredentialProvider,
9+
DatabricksServicePrincipalCredentialProvider,
910
)
1011
from databricks.sql.auth.common import AuthType, ClientContext
1112
from databricks.sql.auth.token_federation import TokenFederationProvider
@@ -15,8 +16,34 @@ def get_auth_provider(cfg: ClientContext, http_client):
1516
# Determine the base auth provider
1617
base_provider: Optional[AuthProvider] = None
1718

19+
if cfg.auth_type and cfg.auth_type not in [t.value for t in AuthType]:
20+
# Never fall through to the interactive browser login below for an
21+
# auth_type this connector does not implement (e.g. a typo).
22+
raise ValueError(
23+
f"Unsupported auth_type {cfg.auth_type!r}; supported: "
24+
+ ", ".join(t.value for t in AuthType)
25+
)
26+
1827
if cfg.credentials_provider:
1928
base_provider = ExternalAuthProvider(cfg.credentials_provider)
29+
elif cfg.oauth_client_secret and cfg.auth_type != AuthType.AZURE_SP_M2M.value:
30+
if cfg.auth_type in [
31+
AuthType.DATABRICKS_OAUTH.value,
32+
AuthType.AZURE_OAUTH.value,
33+
]:
34+
raise ValueError(
35+
f"auth_type={cfg.auth_type!r} selects the interactive (U2M) flow, "
36+
"but oauth_client_secret was also provided (M2M). Drop "
37+
"oauth_client_secret for U2M, or drop auth_type for M2M."
38+
)
39+
base_provider = ExternalAuthProvider(
40+
DatabricksServicePrincipalCredentialProvider(
41+
cfg.hostname,
42+
cfg.oauth_client_id,
43+
cfg.oauth_client_secret,
44+
http_client,
45+
)
46+
)
2047
elif cfg.auth_type == AuthType.AZURE_SP_M2M.value:
2148
base_provider = ExternalAuthProvider(
2249
AzureServicePrincipalCredentialProvider(
@@ -132,5 +159,8 @@ def get_python_sql_connector_auth_provider(hostname: str, http_client, **kwargs)
132159
oauth_persistence=kwargs.get("experimental_oauth_persistence"),
133160
credentials_provider=kwargs.get("credentials_provider"),
134161
identity_federation_client_id=kwargs.get("identity_federation_client_id"),
162+
oauth_client_secret=kwargs.get("oauth_client_secret"),
135163
)
164+
if cfg.oauth_client_secret and not kwargs.get("oauth_client_id"):
165+
raise ValueError("OAuth M2M needs oauth_client_id with oauth_client_secret")
136166
return get_auth_provider(cfg, http_client)

‎src/databricks/sql/auth/authenticators.py‎

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -236,3 +236,47 @@ def header_factory() -> Dict[str, str]:
236236
return headers
237237

238238
return header_factory
239+
240+
241+
class DatabricksServicePrincipalCredentialProvider(CredentialsProvider):
242+
"""
243+
OAuth machine-to-machine (client credentials) authentication for a
244+
Databricks service principal.
245+
246+
Tokens come from the workspace's ``/oidc/v1/token`` endpoint and are
247+
refreshed by the token source when they expire, so a long-lived connection
248+
keeps working past the token lifetime.
249+
250+
Attributes:
251+
hostname (str): The normalized workspace URL (``https://<host>/``).
252+
client_id (str): The service principal's OAuth client (application) ID.
253+
client_secret (str): The service principal's OAuth secret.
254+
"""
255+
256+
DEFAULT_SCOPE = "all-apis"
257+
258+
def __init__(self, hostname, client_id, client_secret, http_client):
259+
self.hostname = hostname
260+
self.client_id = client_id
261+
self.client_secret = client_secret
262+
self._http_client = http_client
263+
264+
def auth_type(self) -> str:
265+
return "oauth-m2m"
266+
267+
def __call__(self, *args, **kwargs) -> HeaderFactory:
268+
source = ClientCredentialsTokenSource(
269+
token_url=f"{self.hostname.rstrip('/')}/oidc/v1/token",
270+
client_id=self.client_id,
271+
client_secret=self.client_secret,
272+
http_client=self._http_client,
273+
extra_params={"scope": self.DEFAULT_SCOPE},
274+
)
275+
276+
def header_factory() -> Dict[str, str]:
277+
token = source.get_token()
278+
return {
279+
HttpHeader.AUTHORIZATION.value: f"{token.token_type} {token.access_token}"
280+
}
281+
282+
return header_factory

‎src/databricks/sql/auth/common.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ def __init__(
3838
oauth_persistence=None,
3939
credentials_provider=None,
4040
identity_federation_client_id: Optional[str] = None,
41+
oauth_client_secret: Optional[str] = None,
4142
# HTTP client configuration parameters
4243
ssl_options=None, # SSLOptions type
4344
socket_timeout: Optional[float] = None,
@@ -59,6 +60,7 @@ def __init__(
5960
self.auth_type = auth_type
6061
self.oauth_scopes = oauth_scopes
6162
self.oauth_client_id = oauth_client_id
63+
self.oauth_client_secret = oauth_client_secret
6264
self.azure_client_id = azure_client_id
6365
self.azure_client_secret = azure_client_secret
6466
self.azure_tenant_id = azure_tenant_id

‎src/databricks/sql/auth/oauth.py‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import base64
2+
import dataclasses
23
import hashlib
34
import json
45
import logging
@@ -341,7 +342,12 @@ def refresh(self) -> Token:
341342
method=HttpMethod.POST, url=self.token_url, headers=headers, body=data
342343
)
343344
if response.status == 200:
344-
oauth_response = OAuthResponse(**json.loads(response.data.decode("utf-8")))
345+
payload = json.loads(response.data.decode("utf-8"))
346+
# Token endpoints add fields such as ``scope``; keep the known ones.
347+
known = {f.name for f in dataclasses.fields(OAuthResponse)}
348+
oauth_response = OAuthResponse(
349+
**{k: v for k, v in payload.items() if k in known}
350+
)
345351
return Token(
346352
oauth_response.access_token,
347353
oauth_response.token_type,

‎src/databricks/sql/client.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -230,6 +230,12 @@ def __init__(
230230
oauth_client_id: `str`, optional
231231
custom oauth client_id. If not specified, it will use the built-in client_id of databricks-sql-python.
232232
233+
oauth_client_secret: `str`, optional
234+
OAuth secret of a service principal. Together with `oauth_client_id`
235+
(the service principal's client ID) this selects OAuth
236+
machine-to-machine authentication (client credentials, scope
237+
`all-apis`), with tokens refreshed as they expire.
238+
233239
oauth_redirect_port: `int`, optional
234240
port of the oauth redirect uri (localhost). This is required when custom oauth client_id
235241
`oauth_client_id` is set

‎tests/unit/test_auth.py‎

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import unittest
2+
from urllib.parse import parse_qs
23
import pytest
34
from unittest.mock import patch, MagicMock
45
import jwt
@@ -152,6 +153,69 @@ def test_get_python_sql_connector_auth_provider_access_token(self):
152153
auth_provider.add_headers(headers)
153154
self.assertEqual(headers["Authorization"], "Bearer dpi123")
154155

156+
def test_get_python_sql_connector_auth_provider_oauth_m2m(self):
157+
"""oauth_client_id + oauth_client_secret authenticate with client
158+
credentials instead of starting an interactive login."""
159+
http_client = MagicMock()
160+
http_client.request.return_value = MagicMock(
161+
status=200,
162+
data=json.dumps(
163+
{
164+
"access_token": "m2m-token",
165+
"token_type": "Bearer",
166+
"expires_in": 3600,
167+
"scope": "all-apis",
168+
}
169+
).encode(),
170+
)
171+
with patch(
172+
"databricks.sql.auth.auth.DatabricksOAuthProvider",
173+
side_effect=AssertionError("interactive login started"),
174+
), patch("databricks.sql.auth.oauth.Token.is_expired", return_value=False):
175+
auth_provider = get_python_sql_connector_auth_provider(
176+
"example.cloud.databricks.com",
177+
http_client,
178+
oauth_client_id="sp-client-id",
179+
oauth_client_secret="sp-secret",
180+
)
181+
headers = {}
182+
auth_provider.add_headers(headers)
183+
self.assertEqual(headers["Authorization"], "Bearer m2m-token")
184+
request = http_client.request.call_args.kwargs
185+
self.assertEqual(
186+
request["url"], "https://example.cloud.databricks.com/oidc/v1/token"
187+
)
188+
body = parse_qs(request["body"])
189+
self.assertEqual(body["grant_type"], ["client_credentials"])
190+
self.assertEqual(body["client_id"], ["sp-client-id"])
191+
self.assertEqual(body["client_secret"], ["sp-secret"])
192+
self.assertEqual(body["scope"], ["all-apis"])
193+
194+
def test_get_python_sql_connector_auth_provider_oauth_m2m_errors(self):
195+
with self.assertRaisesRegex(ValueError, "needs oauth_client_id"):
196+
get_python_sql_connector_auth_provider(
197+
"example.cloud.databricks.com", MagicMock(), oauth_client_secret="s"
198+
)
199+
with self.assertRaisesRegex(ValueError, "interactive"):
200+
get_python_sql_connector_auth_provider(
201+
"example.cloud.databricks.com",
202+
MagicMock(),
203+
auth_type="databricks-oauth",
204+
oauth_client_id="c",
205+
oauth_client_secret="s",
206+
)
207+
208+
def test_get_python_sql_connector_auth_provider_unknown_auth_type(self):
209+
"""An unsupported auth_type must not fall back to a browser login."""
210+
with patch(
211+
"databricks.sql.auth.auth.DatabricksOAuthProvider",
212+
side_effect=AssertionError("interactive login started"),
213+
):
214+
with self.assertRaisesRegex(ValueError, "Unsupported auth_type"):
215+
get_python_sql_connector_auth_provider(
216+
"example.cloud.databricks.com", MagicMock(), auth_type="oauth-u2m"
217+
)
218+
155219
def test_get_python_sql_connector_auth_provider_external(self):
156220
class MyProvider(CredentialsProvider):
157221
def auth_type(self) -> str:

0 commit comments

Comments
 (0)