|
1 | 1 | import unittest |
| 2 | +from urllib.parse import parse_qs |
2 | 3 | import pytest |
3 | 4 | from unittest.mock import patch, MagicMock |
4 | 5 | import jwt |
@@ -152,6 +153,69 @@ def test_get_python_sql_connector_auth_provider_access_token(self): |
152 | 153 | auth_provider.add_headers(headers) |
153 | 154 | self.assertEqual(headers["Authorization"], "Bearer dpi123") |
154 | 155 |
|
| 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 | + |
155 | 219 | def test_get_python_sql_connector_auth_provider_external(self): |
156 | 220 | class MyProvider(CredentialsProvider): |
157 | 221 | def auth_type(self) -> str: |
|
0 commit comments