FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Use GcpTarget in KaggleKernelCredentials · AI-For-Rural/docker-python@17c022a · GitHub

Commit 17c022a

Browse files
committed
Use GcpTarget in KaggleKernelCredentials
1 parent c396a06 commit 17c022a

3 files changed

Lines changed: 32 additions & 10 deletions

File tree

‎patches/kaggle_gcp.py‎

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
from google.cloud import bigquery
55
from google.cloud.exceptions import Forbidden
66
from google.cloud.bigquery._http import Connection
7-
from kaggle_secrets import UserSecretsClient
7+
from kaggle_secrets import GcpTarget, UserSecretsClient
88

99
from log import Log
1010

@@ -29,26 +29,35 @@ def add_integration(self, integration_name):
2929
def has_bigquery(self):
3030
return 'bigquery' in self.integrations.keys()
3131

32+
def has_gcs(self):
33+
return 'gcs' in self.integrations.keys()
34+
3235

3336
class KaggleKernelCredentials(credentials.Credentials):
3437
"""Custom Credentials used to authenticate using the Kernel's connected OAuth account.
3538
Example usage:
3639
client = bigquery.Client(project='ANOTHER_PROJECT',
3740
credentials=KaggleKernelCredentials())
3841
"""
42+
def __init__(self, target=GcpTarget.BIGQUERY):
43+
super().__init__()
44+
self.target = target
3945

4046
def refresh(self, request):
4147
try:
4248
client = UserSecretsClient()
43-
self.token, self.expiry = client.get_bigquery_access_token()
49+
if self.target == GcpTarget.BIGQUERY:
50+
self.token, self.expiry = client.get_bigquery_access_token()
51+
elif self.target == GcpTarget.GCS:
52+
self.token, self.expiry = client._get_gcs_access_token()
4453
except ConnectionError as e:
4554
Log.error(f"Connection error trying to refresh access token: {e}")
4655
print("There was a connection error trying to fetch the access token. "
47-
"Please ensure internet is on in order to use the BigQuery Integration.")
56+
f"Please ensure internet is on in order to use the {self.target.service} Integration.")
4857
raise RefreshError('Unable to refresh access token due to connection error.') from e
4958
except Exception as e:
5059
Log.error(f"Error trying to refresh access token: {e}")
51-
if (not get_integrations().has_bigquery()):
60+
if (not get_integrations().has_bigquery() and self.target == GcpTarget.BIGQUERY):
5261
Log.error(f"No bigquery integration found.")
5362
print(
5463
'Please ensure you have selected a BigQuery account in the Kernels Settings sidebar.')

‎patches/kaggle_secrets.py‎

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -28,8 +28,21 @@ class BackendError(Exception):
2828

2929
@unique
3030
class GcpTarget(Enum):
31-
BIGQUERY = 1
32-
GCS = 2
31+
"""Enum class to store GCP targets."""
32+
BIGQUERY = (1, "BigQuery")
33+
GCS = (2, "Google Cloud Storage")
34+
35+
def __init__(self, target, service):
36+
self._target = target
37+
self._service = service
38+
39+
@property
40+
def target(self):
41+
return self._target
42+
43+
@property
44+
def service(self):
45+
return self._service
3346

3447

3548
class UserSecretsClient():
@@ -91,7 +104,7 @@ def _get_gcs_access_token(self) -> Tuple[str, Optional[datetime]]:
91104

92105
def _get_access_token(self, target: GcpTarget) -> Tuple[str, Optional[datetime]]:
93106
request_body = {
94-
'Target': target.value
107+
'Target': target.target
95108
}
96109
response_json = self._make_post_request(request_body)
97110
if 'secret' not in response_json:

‎tests/test_user_secrets.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -99,10 +99,10 @@ def call_get_gcs_access_token():
9999
secret_response = client._get_gcs_access_token()
100100
self.assertEqual(secret_response, (secret, now + timedelta(seconds=3600)))
101101
self._test_client(call_get_bigquery_access_token,
102-
'/requests/GetUserSecretRequest', {'Target': GcpTarget.BIGQUERY.value, 'JWE': _TEST_JWT},
102+
'/requests/GetUserSecretRequest', {'Target': GcpTarget.BIGQUERY.target, 'JWE': _TEST_JWT},
103103
secret=secret)
104104
self._test_client(call_get_gcs_access_token,
105-
'/requests/GetUserSecretRequest', {'Target': GcpTarget.GCS.value, 'JWE': _TEST_JWT},
105+
'/requests/GetUserSecretRequest', {'Target': GcpTarget.GCS.target, 'JWE': _TEST_JWT},
106106
secret=secret)
107107

108108
def test_get_access_token_handles_unsuccessful(self):
@@ -111,4 +111,4 @@ def call_get_access_token():
111111
with self.assertRaises(BackendError):
112112
client.get_bigquery_access_token()
113113
self._test_client(call_get_access_token,
114-
'/requests/GetUserSecretRequest', {'Target': GcpTarget.BIGQUERY.value, 'JWE': _TEST_JWT}, success=False)
114+
'/requests/GetUserSecretRequest', {'Target': GcpTarget.BIGQUERY.target, 'JWE': _TEST_JWT}, success=False)

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL