| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,7 +4,7 @@ | |||
| 4 | 4 | from google.cloud import bigquery | |
| 5 | 5 | from google.cloud.exceptions import Forbidden | |
| 6 | 6 | from google.cloud.bigquery._http import Connection | |
| 7 | - from kaggle_secrets import UserSecretsClient | ||
| 7 | + from kaggle_secrets import GcpTarget, UserSecretsClient | ||
| 8 | 8 | ||
| 9 | 9 | from log import Log | |
| 10 | 10 | ||
@@ -29,26 +29,35 @@ def add_integration(self, integration_name): | |||
| 29 | 29 | def has_bigquery(self): | |
| 30 | 30 | return 'bigquery' in self.integrations.keys() | |
| 31 | 31 | ||
| 32 | + def has_gcs(self): | ||
| 33 | + return 'gcs' in self.integrations.keys() | ||
| 34 | + | ||
| 32 | 35 | ||
| 33 | 36 | class KaggleKernelCredentials(credentials.Credentials): | |
| 34 | 37 | """Custom Credentials used to authenticate using the Kernel's connected OAuth account. | |
| 35 | 38 | Example usage: | |
| 36 | 39 | client = bigquery.Client(project='ANOTHER_PROJECT', | |
| 37 | 40 | credentials=KaggleKernelCredentials()) | |
| 38 | 41 | """ | |
| 42 | + def __init__(self, target=GcpTarget.BIGQUERY): | ||
| 43 | + super().__init__() | ||
| 44 | + self.target = target | ||
| 39 | 45 | ||
| 40 | 46 | def refresh(self, request): | |
| 41 | 47 | try: | |
| 42 | 48 | 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() | ||
| 44 | 53 | except ConnectionError as e: | |
| 45 | 54 | Log.error(f"Connection error trying to refresh access token: {e}") | |
| 46 | 55 | 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.") | ||
| 48 | 57 | raise RefreshError('Unable to refresh access token due to connection error.') from e | |
| 49 | 58 | except Exception as e: | |
| 50 | 59 | 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): | ||
| 52 | 61 | Log.error(f"No bigquery integration found.") | |
| 53 | 62 | print( | |
| 54 | 63 | 'Please ensure you have selected a BigQuery account in the Kernels Settings sidebar.') | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -28,8 +28,21 @@ class BackendError(Exception): | |||
| 28 | 28 | ||
| 29 | 29 | @unique | |
| 30 | 30 | 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 | ||
| 33 | 46 | ||
| 34 | 47 | ||
| 35 | 48 | class UserSecretsClient(): | |
@@ -91,7 +104,7 @@ def _get_gcs_access_token(self) -> Tuple[str, Optional[datetime]]: | |||
| 91 | 104 | ||
| 92 | 105 | def _get_access_token(self, target: GcpTarget) -> Tuple[str, Optional[datetime]]: | |
| 93 | 106 | request_body = { | |
| 94 | - 'Target': target.value | ||
| 107 | + 'Target': target.target | ||
| 95 | 108 | } | |
| 96 | 109 | response_json = self._make_post_request(request_body) | |
| 97 | 110 | if 'secret' not in response_json: | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -99,10 +99,10 @@ def call_get_gcs_access_token(): | |||
| 99 | 99 | secret_response = client._get_gcs_access_token() | |
| 100 | 100 | self.assertEqual(secret_response, (secret, now + timedelta(seconds=3600))) | |
| 101 | 101 | 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}, | ||
| 103 | 103 | secret=secret) | |
| 104 | 104 | 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}, | ||
| 106 | 106 | secret=secret) | |
| 107 | 107 | ||
| 108 | 108 | def test_get_access_token_handles_unsuccessful(self): | |
@@ -111,4 +111,4 @@ def call_get_access_token(): | |||
| 111 | 111 | with self.assertRaises(BackendError): | |
| 112 | 112 | client.get_bigquery_access_token() | |
| 113 | 113 | 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) | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments