99from typing import List
1010
1111BASE_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
69166def 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 )
0 commit comments