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

Add a Kaggle dataset client to support getting GCS paths for TPU and … · HpcDataLab/docker-python@59eaede · GitHub

Commit 59eaede

Browse files
committed
Add a Kaggle dataset client to support getting GCS paths for TPU and AutoML
1 parent 0f59abb commit 59eaede

3 files changed

Lines changed: 212 additions & 0 deletions

File tree

‎patches/kaggle_datasets.py‎

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff 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']

‎patches/kaggle_web_client.py‎

Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff 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

‎tests/test_datasets.py‎

Lines changed: 130 additions & 0 deletions
Original file line numberDiff line numberDiff 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)

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL