| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 0f59abb commit 59eaede
3 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,24 @@ | |||
| 1 | + import os | ||
| 2 | + from kaggle_web_client import KaggleWebClient | ||
| 3 | + | ||
| 4 | + _KAGGLE_TPU_NAME_ENV_VAR_NAME = 'TPU_NAME' | ||
| 5 | + | ||
| 6 | + class KaggleDatasets: | ||
| 7 | + GET_GCS_PATH_ENDPOINT = '/requests/CopyDatasetVersionToKnownGcsBucketRequest' | ||
| 8 | + | ||
| 9 | + # Integration types for GCS | ||
| 10 | + AUTO_ML = 1 | ||
| 11 | + TPU = 2 | ||
| 12 | + | ||
| 13 | + def __init__(self): | ||
| 14 | + self.web_client = KaggleWebClient() | ||
| 15 | + self.has_tpu = os.getenv(_KAGGLE_TPU_NAME_ENV_VAR_NAME) is not None | ||
| 16 | + | ||
| 17 | + def get_gcs_path(self, dataset_dir: str = None) -> str: | ||
| 18 | + integration_type = self.TPU if self.has_tpu else self.AUTO_ML | ||
| 19 | + data = { | ||
| 20 | + 'MountSlug': dataset_dir, | ||
| 21 | + 'IntegrationType': integration_type, | ||
| 22 | + } | ||
| 23 | + result = self.web_client.make_post_request(data, self.GET_GCS_PATH_ENDPOINT) | ||
| 24 | + return result['destinationBucket'] | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,58 @@ | |||
| 1 | + | ||
| 2 | + import json | ||
| 3 | + import os | ||
| 4 | + import socket | ||
| 5 | + import urllib.request | ||
| 6 | + from datetime import datetime, timedelta | ||
| 7 | + from enum import Enum, unique | ||
| 8 | + from typing import Optional, Tuple | ||
| 9 | + from urllib.error import HTTPError, URLError | ||
| 10 | + from kaggle_secrets import (_KAGGLE_DEFAULT_URL_BASE, | ||
| 11 | + _KAGGLE_URL_BASE_ENV_VAR_NAME, | ||
| 12 | + _KAGGLE_USER_SECRETS_TOKEN_ENV_VAR_NAME, | ||
| 13 | + CredentialError, BackendError, ValidationError) | ||
| 14 | + | ||
| 15 | + class KaggleWebClient: | ||
| 16 | + TIMEOUT_SECS = 90 | ||
| 17 | + | ||
| 18 | + def __init__(self): | ||
| 19 | + url_base_override = os.getenv(_KAGGLE_URL_BASE_ENV_VAR_NAME) | ||
| 20 | + self.url_base = url_base_override or _KAGGLE_DEFAULT_URL_BASE | ||
| 21 | + # Follow the OAuth 2.0 Authorization standard (https://tools.ietf.org/html/rfc6750) | ||
| 22 | + self.jwt_token = os.getenv(_KAGGLE_USER_SECRETS_TOKEN_ENV_VAR_NAME) | ||
| 23 | + if self.jwt_token is None: | ||
| 24 | + raise CredentialError( | ||
| 25 | + 'A JWT Token is required to call Kaggle, ' | ||
| 26 | + f'but none found in environment variable {_KAGGLE_USER_SECRETS_TOKEN_ENV_VAR_NAME}') | ||
| 27 | + self.headers = { | ||
| 28 | + 'Content-type': 'application/json', | ||
| 29 | + 'Authorization': f'Bearer {self.jwt_token}', | ||
| 30 | + } | ||
| 31 | + | ||
| 32 | + def make_post_request(self, data: dict, endpoint: str) -> dict: | ||
| 33 | + url = f'{self.url_base}{endpoint}' | ||
| 34 | + request_body = dict(data) | ||
| 35 | + request_body['JWE'] = self.jwt_token | ||
| 36 | + req = urllib.request.Request(url, headers=self.headers, data=bytes( | ||
| 37 | + json.dumps(request_body), encoding="utf-8")) | ||
| 38 | + try: | ||
| 39 | + with urllib.request.urlopen(req, timeout=self.TIMEOUT_SECS) as response: | ||
| 40 | + response_json = json.loads(response.read()) | ||
| 41 | + if not response_json.get('wasSuccessful') or 'result' not in response_json: | ||
| 42 | + raise BackendError( | ||
| 43 | + f'Unexpected response from the service. Response: {response_json}.') | ||
| 44 | + return response_json['result'] | ||
| 45 | + except (URLError, socket.timeout) as e: | ||
| 46 | + if isinstance( | ||
| 47 | + e, socket.timeout) or isinstance( | ||
| 48 | + e.reason, socket.timeout): | ||
| 49 | + raise ConnectionError( | ||
| 50 | + 'Timeout error trying to communicate with service. Please ensure internet is on.') from e | ||
| 51 | + raise ConnectionError( | ||
| 52 | + 'Connection error trying to communicate with service.') from e | ||
| 53 | + except HTTPError as e: | ||
| 54 | + if e.code == 401 or e.code == 403: | ||
| 55 | + raise CredentialError( | ||
| 56 | + f'Service responded with error code {e.code}.' | ||
| 57 | + ' Please ensure you have access to the resource.') from e | ||
| 58 | + raise BackendError('Unexpected response from the service.') from e | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,130 @@ | |||
| 1 | + import json | ||
| 2 | + import os | ||
| 3 | + import threading | ||
| 4 | + import unittest | ||
| 5 | + from http.server import BaseHTTPRequestHandler, HTTPServer | ||
| 6 | + from test.support import EnvironmentVarGuard | ||
| 7 | + from urllib.parse import urlparse | ||
| 8 | + from datetime import datetime, timedelta | ||
| 9 | + import mock | ||
| 10 | + | ||
| 11 | + from google.auth.exceptions import DefaultCredentialsError | ||
| 12 | + from google.cloud import bigquery | ||
| 13 | + from kaggle_secrets import (_KAGGLE_URL_BASE_ENV_VAR_NAME, | ||
| 14 | + _KAGGLE_USER_SECRETS_TOKEN_ENV_VAR_NAME, | ||
| 15 | + CredentialError, BackendError, ValidationError) | ||
| 16 | + from kaggle_web_client import KaggleWebClient | ||
| 17 | + from kaggle_datasets import KaggleDatasets, _KAGGLE_TPU_NAME_ENV_VAR_NAME | ||
| 18 | + | ||
| 19 | + _TEST_JWT = 'test-secrets-key' | ||
| 20 | + | ||
| 21 | + _TPU_GCS_BUCKET = 'gs://kds-tpu-ea1971a458ffd4cd51389e7574c022ecc0a82bb1b52ccef08c8a' | ||
| 22 | + _AUTOML_GCS_BUCKET = 'gs://kds-automl-ea1971a458ffd4cd51389e7574c022ecc0a82bb1b52ccef08c8a' | ||
| 23 | + | ||
| 24 | + class GcsDatasetsHTTPHandler(BaseHTTPRequestHandler): | ||
| 25 | + | ||
| 26 | + def set_request(self): | ||
| 27 | + raise NotImplementedError() | ||
| 28 | + | ||
| 29 | + def get_response(self): | ||
| 30 | + raise NotImplementedError() | ||
| 31 | + | ||
| 32 | + def do_HEAD(s): | ||
| 33 | + s.send_response(200) | ||
| 34 | + | ||
| 35 | + def do_POST(s): | ||
| 36 | + s.set_request() | ||
| 37 | + s.send_response(200) | ||
| 38 | + s.send_header("Content-type", "application/json") | ||
| 39 | + s.end_headers() | ||
| 40 | + s.wfile.write(json.dumps(s.get_response()).encode("utf-8")) | ||
| 41 | + | ||
| 42 | + | ||
| 43 | + class TestDatasets(unittest.TestCase): | ||
| 44 | + SERVER_ADDRESS = urlparse(os.getenv(_KAGGLE_URL_BASE_ENV_VAR_NAME, default="http://127.0.0.1:8001")) | ||
| 45 | + | ||
| 46 | + def _test_client(self, client_func, expected_path, expected_body, is_tpu=True, success=True): | ||
| 47 | + _request = {} | ||
| 48 | + | ||
| 49 | + class GetGcsPathHandler(GcsDatasetsHTTPHandler): | ||
| 50 | + | ||
| 51 | + def set_request(self): | ||
| 52 | + _request['path'] = self.path | ||
| 53 | + content_len = int(self.headers.get('Content-Length')) | ||
| 54 | + _request['body'] = json.loads(self.rfile.read(content_len)) | ||
| 55 | + _request['headers'] = self.headers | ||
| 56 | + | ||
| 57 | + def get_response(self): | ||
| 58 | + if success: | ||
| 59 | + gcs_path = _TPU_GCS_BUCKET if is_tpu else _AUTOML_GCS_BUCKET | ||
| 60 | + return {'result': { | ||
| 61 | + 'destinationBucket': gcs_path, | ||
| 62 | + 'destinationPath': None}, 'wasSuccessful': "true"} | ||
| 63 | + else: | ||
| 64 | + return {'wasSuccessful': "false"} | ||
| 65 | + | ||
| 66 | + env = EnvironmentVarGuard() | ||
| 67 | + env.set(_KAGGLE_USER_SECRETS_TOKEN_ENV_VAR_NAME, _TEST_JWT) | ||
| 68 | + if is_tpu: | ||
| 69 | + env.set(_KAGGLE_TPU_NAME_ENV_VAR_NAME, 'FAKE_TPU') | ||
| 70 | + with env: | ||
| 71 | + with HTTPServer((self.SERVER_ADDRESS.hostname, self.SERVER_ADDRESS.port), GetGcsPathHandler) as httpd: | ||
| 72 | + threading.Thread(target=httpd.serve_forever).start() | ||
| 73 | + | ||
| 74 | + try: | ||
| 75 | + client_func() | ||
| 76 | + finally: | ||
| 77 | + httpd.shutdown() | ||
| 78 | + | ||
| 79 | + path, headers, body = _request['path'], _request['headers'], _request['body'] | ||
| 80 | + self.assertEqual( | ||
| 81 | + path, | ||
| 82 | + expected_path, | ||
| 83 | + msg="Fake server did not receive the right request from the KaggleDatasets client.") | ||
| 84 | + self.assertEqual( | ||
| 85 | + body, | ||
| 86 | + expected_body, | ||
| 87 | + msg="Fake server did not receive the right body from the KaggleDatasets client.") | ||
| 88 | + self.assertIn('Content-Type', headers.keys(), | ||
| 89 | + msg="Fake server did not receive a Content-Type header from the KaggleDatasets client.") | ||
| 90 | + self.assertEqual('application/json', headers.get('Content-Type'), | ||
| 91 | + msg="Fake server did not receive an application/json content type header from the KaggleDatasets client.") | ||
| 92 | + self.assertIn('Authorization', headers.keys(), | ||
| 93 | + msg="Fake server did not receive an Authorization header from the KaggleDatasets client.") | ||
| 94 | + self.assertEqual(f'Bearer {_TEST_JWT}', headers.get('Authorization'), | ||
| 95 | + msg="Fake server did not receive the right Authorization header from the KaggleDatasets client.") | ||
| 96 | + | ||
| 97 | + def test_no_token_fails(self): | ||
| 98 | + env = EnvironmentVarGuard() | ||
| 99 | + env.unset(_KAGGLE_USER_SECRETS_TOKEN_ENV_VAR_NAME) | ||
| 100 | + with env: | ||
| 101 | + with self.assertRaises(CredentialError): | ||
| 102 | + client = KaggleDatasets() | ||
| 103 | + | ||
| 104 | + def test_get_gcs_path_tpu_succeeds(self): | ||
| 105 | + def call_get_gcs_path(): | ||
| 106 | + client = KaggleDatasets() | ||
| 107 | + gcs_path = client.get_gcs_path() | ||
| 108 | + self.assertEqual(gcs_path, _TPU_GCS_BUCKET) | ||
| 109 | + self._test_client(call_get_gcs_path, | ||
| 110 | + '/requests/CopyDatasetVersionToKnownGcsBucketRequest', {'MountSlug': None, 'IntegrationType': 2, 'JWE': _TEST_JWT}, | ||
| 111 | + is_tpu=True) | ||
| 112 | + | ||
| 113 | + def test_get_gcs_path_automl_succeeds(self): | ||
| 114 | + def call_get_gcs_path(): | ||
| 115 | + client = KaggleDatasets() | ||
| 116 | + gcs_path = client.get_gcs_path() | ||
| 117 | + self.assertEqual(gcs_path, _AUTOML_GCS_BUCKET) | ||
| 118 | + self._test_client(call_get_gcs_path, | ||
| 119 | + '/requests/CopyDatasetVersionToKnownGcsBucketRequest', {'MountSlug': None, 'IntegrationType': 1, 'JWE': _TEST_JWT}, | ||
| 120 | + is_tpu=False) | ||
| 121 | + | ||
| 122 | + def test_get_gcs_path_handles_unsuccessful(self): | ||
| 123 | + def call_get_gcs_path(): | ||
| 124 | + client = KaggleDatasets() | ||
| 125 | + with self.assertRaises(BackendError): | ||
| 126 | + gcs_path = client.get_gcs_path() | ||
| 127 | + self._test_client(call_get_gcs_path, | ||
| 128 | + '/requests/CopyDatasetVersionToKnownGcsBucketRequest', {'MountSlug': None, 'IntegrationType': 2, 'JWE': _TEST_JWT}, | ||
| 129 | + is_tpu=True, | ||
| 130 | + success=False) | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments