Skip to content

Commit 54d13bf

Browse files
committed
feat(downloader): Support pulling Toolkit bundle from OCI registry
1 parent bbf9396 commit 54d13bf

3 files changed

Lines changed: 220 additions & 23 deletions

File tree

dockerfiles/cache-downloader/Dockerfile

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
FROM alpine:3.22
22

3-
RUN apk add --no-cache bash aws-cli python3
3+
RUN apk add --no-cache bash aws-cli ca-certificates python3 zstd
4+
5+
COPY --from=docker.io/regclient/regctl:v0.11.5 /regctl /usr/local/bin/regctl
46

57
COPY cache-downloader.py /usr/bin/cache-downloader
68
RUN chmod 755 /usr/bin/cache-downloader
@@ -15,4 +17,4 @@ RUN if [ -z "$PYTHON_VERSIONS" ]; then echo "PYTHON_VERSIONS argument not provid
1517

1618
ENV PYTHONUNBUFFERED=1
1719

18-
CMD [ "cache-downloader" ]
20+
CMD [ "cache-downloader" ]

dockerfiles/cache-downloader/cache-downloader.py

Lines changed: 134 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -9,29 +9,34 @@
99
from typing import List
1010

1111
BASE_PATH = "/deepnote-toolkit/"
12+
DEFAULT_DOWNLOAD_METHOD = "s3"
13+
OCI_ZSTD_DOWNLOAD_METHOD = "oci-zstd"
14+
DEFAULT_TOOLKIT_BUNDLE_OCI_REPOSITORY = "docker.io/deepnote/toolkit-bundle"
1215

1316

14-
def download_dependency(
15-
release_name: str, python_version: str, toolkit_index_bucket_name: str
16-
):
17-
"""Download the dependencies for the given Python version and release name."""
17+
def build_toolkit_bundle_ref(
18+
release_name: str,
19+
python_version: str,
20+
oci_repository: str = DEFAULT_TOOLKIT_BUNDLE_OCI_REPOSITORY,
21+
) -> str:
22+
return f"{oci_repository}:{release_name}-python{python_version}-tar-zst"
1823

19-
version_path = os.path.join(BASE_PATH, release_name, f"python{python_version}")
20-
done_file = os.path.join(version_path, f"{python_version}-done")
2124

22-
if Path(done_file).is_file():
23-
print(
24-
f"{datetime.datetime.now()}: {release_name} python{python_version} already cached, skipping download"
25-
)
26-
return
25+
def _read_stderr(process: subprocess.Popen) -> bytes:
26+
if process.stderr is None:
27+
return b""
28+
return process.stderr.read()
2729

28-
# Create the version directory if it doesn't exist
29-
os.makedirs(version_path, exist_ok=True)
3030

31+
def download_dependency_from_s3(
32+
release_name: str,
33+
python_version: str,
34+
toolkit_index_bucket_name: str,
35+
version_path: str,
36+
) -> None:
3137
s3_path = f"s3://{toolkit_index_bucket_name}/deepnote-toolkit/{release_name}/python{python_version}.tar"
3238
print(f"{datetime.datetime.now()}: Downloading {release_name} {s3_path}")
3339

34-
# Use Popen to stream the data
3540
aws_process = subprocess.Popen(
3641
["aws", "s3", "cp", "--no-sign-request", s3_path, "-"],
3742
stdout=subprocess.PIPE,
@@ -44,16 +49,13 @@ def download_dependency(
4449
stderr=subprocess.PIPE,
4550
)
4651
aws_process.stdout.close() # Allow aws_process to receive a SIGPIPE if tar_process exits.
47-
_, tar_process_stderr = (
48-
tar_process.communicate()
49-
) # Wait for tar_process to complete
52+
_, tar_process_stderr = tar_process.communicate()
5053
aws_process_returncode = aws_process.wait()
5154

5255
if aws_process_returncode != 0:
5356
raise Exception(
54-
f"Error downloading {release_name} (aws s3 command failed): aws stderr: {aws_process.stderr.read()}, tar stderr: {tar_process_stderr}"
57+
f"Error downloading {release_name} (aws s3 command failed): aws stderr: {_read_stderr(aws_process)}, tar stderr: {tar_process_stderr}"
5558
)
56-
# Check for errors
5759
if tar_process.returncode != 0:
5860
raise Exception(
5961
f"Error downloading {release_name} (tar command failed): {tar_process_stderr}"
@@ -62,12 +64,111 @@ def download_dependency(
6264
print(
6365
f"{datetime.datetime.now()}: Done downloading {release_name} {s3_path} and extracting to {version_path}"
6466
)
67+
68+
69+
def download_dependency_from_oci_zstd(
70+
release_name: str,
71+
python_version: str,
72+
oci_repository: str,
73+
version_path: str,
74+
) -> None:
75+
artifact_file = f"python{python_version}.tar.zst"
76+
artifact_ref = build_toolkit_bundle_ref(
77+
release_name, python_version, oci_repository
78+
)
79+
print(f"{datetime.datetime.now()}: Downloading {release_name} {artifact_ref}")
80+
81+
regctl_process = subprocess.Popen(
82+
["regctl", "artifact", "get", "--file", artifact_file, artifact_ref],
83+
stdout=subprocess.PIPE,
84+
stderr=subprocess.PIPE,
85+
)
86+
zstd_process = subprocess.Popen(
87+
["zstd", "-dc"],
88+
stdin=regctl_process.stdout,
89+
stdout=subprocess.PIPE,
90+
stderr=subprocess.PIPE,
91+
)
92+
regctl_process.stdout.close()
93+
tar_process = subprocess.Popen(
94+
["tar", "-xf", "-", "-C", version_path],
95+
stdin=zstd_process.stdout,
96+
stdout=subprocess.PIPE,
97+
stderr=subprocess.PIPE,
98+
)
99+
zstd_process.stdout.close()
100+
101+
_, tar_process_stderr = tar_process.communicate()
102+
zstd_process_returncode = zstd_process.wait()
103+
regctl_process_returncode = regctl_process.wait()
104+
105+
if regctl_process_returncode != 0:
106+
raise Exception(
107+
f"Error downloading {release_name} (regctl command failed): regctl stderr: {_read_stderr(regctl_process)}, zstd stderr: {_read_stderr(zstd_process)}, tar stderr: {tar_process_stderr}"
108+
)
109+
if zstd_process_returncode != 0:
110+
raise Exception(
111+
f"Error downloading {release_name} (zstd command failed): zstd stderr: {_read_stderr(zstd_process)}, tar stderr: {tar_process_stderr}"
112+
)
113+
if tar_process.returncode != 0:
114+
raise Exception(
115+
f"Error downloading {release_name} (tar command failed): {tar_process_stderr}"
116+
)
117+
118+
print(
119+
f"{datetime.datetime.now()}: Done downloading {release_name} {artifact_ref} and extracting to {version_path}"
120+
)
121+
122+
123+
def download_dependency(
124+
release_name: str,
125+
python_version: str,
126+
toolkit_index_bucket_name: str,
127+
download_method: str = DEFAULT_DOWNLOAD_METHOD,
128+
oci_repository: str = DEFAULT_TOOLKIT_BUNDLE_OCI_REPOSITORY,
129+
):
130+
"""Download the dependencies for the given Python version and release name."""
131+
132+
version_path = os.path.join(BASE_PATH, release_name, f"python{python_version}")
133+
done_file = os.path.join(version_path, f"{python_version}-done")
134+
135+
if Path(done_file).is_file():
136+
print(
137+
f"{datetime.datetime.now()}: {release_name} python{python_version} already cached, skipping download"
138+
)
139+
return
140+
141+
os.makedirs(version_path, exist_ok=True)
142+
143+
if download_method == DEFAULT_DOWNLOAD_METHOD:
144+
download_dependency_from_s3(
145+
release_name,
146+
python_version,
147+
toolkit_index_bucket_name,
148+
version_path,
149+
)
150+
elif download_method == OCI_ZSTD_DOWNLOAD_METHOD:
151+
download_dependency_from_oci_zstd(
152+
release_name,
153+
python_version,
154+
oci_repository,
155+
version_path,
156+
)
157+
else:
158+
raise ValueError(
159+
f"Unsupported TOOLKIT_DOWNLOAD_METHOD={download_method!r}; expected one of: {DEFAULT_DOWNLOAD_METHOD}, {OCI_ZSTD_DOWNLOAD_METHOD}"
160+
)
161+
65162
# Create the "done" file
66163
Path(done_file).touch()
67164

68165

69166
def submit_downloading(
70-
python_versions: List[str], release_name: str, toolkit_index_bucket_name: str
167+
python_versions: List[str],
168+
release_name: str,
169+
toolkit_index_bucket_name: str,
170+
download_method: str = DEFAULT_DOWNLOAD_METHOD,
171+
oci_repository: str = DEFAULT_TOOLKIT_BUNDLE_OCI_REPOSITORY,
71172
):
72173
"""Download the dependencies for the given Python versions and release name."""
73174

@@ -78,6 +179,8 @@ def submit_downloading(
78179
release_name,
79180
python_version,
80181
toolkit_index_bucket_name,
182+
download_method,
183+
oci_repository,
81184
)
82185
for python_version in python_versions
83186
]
@@ -151,9 +254,19 @@ def main():
151254
sys.exit(1)
152255

153256
toolkit_index_bucket_name = os.getenv("TOOLKIT_INDEX_BUCKET_NAME")
257+
download_method = os.getenv("TOOLKIT_DOWNLOAD_METHOD", DEFAULT_DOWNLOAD_METHOD)
258+
oci_repository = os.getenv(
259+
"TOOLKIT_BUNDLE_OCI_REPOSITORY", DEFAULT_TOOLKIT_BUNDLE_OCI_REPOSITORY
260+
)
154261

155262
cleanup_old_versions(BASE_PATH, release_name)
156-
submit_downloading(python_versions, release_name, toolkit_index_bucket_name)
263+
submit_downloading(
264+
python_versions,
265+
release_name,
266+
toolkit_index_bucket_name,
267+
download_method,
268+
oci_repository,
269+
)
157270

158271
end_time = datetime.datetime.now()
159272
print("End time:", end_time)

tests/unit/test_cache_downloader.py

Lines changed: 82 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
_spec.loader.exec_module(_mod)
2222

2323
cleanup_old_versions = _mod.cleanup_old_versions
24+
build_toolkit_bundle_ref = _mod.build_toolkit_bundle_ref
2425
download_dependency = _mod.download_dependency
2526

2627

@@ -173,6 +174,14 @@ def _fail_on_v1(path, *args, **kwargs):
173174

174175

175176
class TestDownloadDependency:
177+
def test_build_toolkit_bundle_ref(self):
178+
assert (
179+
build_toolkit_bundle_ref(
180+
"v1.0.0", "3.11", "registry.example.com/deepnote/toolkit-bundle"
181+
)
182+
== "registry.example.com/deepnote/toolkit-bundle:v1.0.0-python3.11-tar-zst"
183+
)
184+
176185
def test_skips_download_when_done_file_exists(self, tmp_path):
177186
"""Already-cached version should not trigger a download."""
178187
release = "v1.0.0"
@@ -214,4 +223,77 @@ def test_downloads_when_done_file_missing(self, tmp_path):
214223
download_dependency(release, py_ver, "fake-bucket")
215224

216225
assert mock_popen.call_count == 2
226+
assert mock_popen.call_args_list[0][0][0] == [
227+
"aws",
228+
"s3",
229+
"cp",
230+
"--no-sign-request",
231+
"s3://fake-bucket/deepnote-toolkit/v1.0.0/python3.11.tar",
232+
"-",
233+
]
234+
assert mock_popen.call_args_list[1][0][0] == [
235+
"tar",
236+
"-xf",
237+
"-",
238+
"-C",
239+
str(version_path),
240+
]
241+
assert (version_path / f"{py_ver}-done").exists()
242+
243+
def test_downloads_from_oci_zstd_when_enabled(self, tmp_path):
244+
"""The OCI zstd feature flag should stream regctl through zstd into tar."""
245+
release = "v1.0.0"
246+
py_ver = "3.11"
247+
version_path = tmp_path / release / f"python{py_ver}"
248+
version_path.mkdir(parents=True)
249+
250+
mock_regctl = MagicMock()
251+
mock_regctl.stdout = MagicMock()
252+
mock_regctl.stderr = MagicMock(read=MagicMock(return_value=b""))
253+
mock_regctl.wait.return_value = 0
254+
255+
mock_zstd = MagicMock()
256+
mock_zstd.stdout = MagicMock()
257+
mock_zstd.stderr = MagicMock(read=MagicMock(return_value=b""))
258+
mock_zstd.wait.return_value = 0
259+
260+
mock_tar = MagicMock()
261+
mock_tar.communicate.return_value = (b"", b"")
262+
mock_tar.returncode = 0
263+
264+
with (
265+
patch.object(_mod, "BASE_PATH", str(tmp_path)),
266+
patch.object(
267+
_mod.subprocess,
268+
"Popen",
269+
side_effect=[mock_regctl, mock_zstd, mock_tar],
270+
) as mock_popen,
271+
):
272+
download_dependency(
273+
release,
274+
py_ver,
275+
"fake-bucket",
276+
download_method="oci-zstd",
277+
oci_repository="registry.example.com/deepnote/toolkit-bundle",
278+
)
279+
280+
assert mock_popen.call_count == 3
281+
assert mock_popen.call_args_list[0][0][0] == [
282+
"regctl",
283+
"artifact",
284+
"get",
285+
"--file",
286+
"python3.11.tar.zst",
287+
"registry.example.com/deepnote/toolkit-bundle:v1.0.0-python3.11-tar-zst",
288+
]
289+
assert mock_popen.call_args_list[1][0][0] == ["zstd", "-dc"]
290+
assert mock_popen.call_args_list[2][0][0] == [
291+
"tar",
292+
"-xf",
293+
"-",
294+
"-C",
295+
str(version_path),
296+
]
297+
mock_regctl.stdout.close.assert_called_once()
298+
mock_zstd.stdout.close.assert_called_once()
217299
assert (version_path / f"{py_ver}-done").exists()

0 commit comments

Comments
 (0)