|
1 | 1 | from typing import Iterable |
| 2 | +import json |
| 3 | +import os |
2 | 4 | import requests |
3 | 5 | from ebrains_drive.exceptions import DoesNotExist, InvalidParameter, UpstreamAPIException |
4 | 6 | from ebrains_drive.files import DataproxyFile |
5 | 7 | from ebrains_drive.utils import on_401_raise_unauthorized |
6 | 8 | from io import IOBase |
| 9 | +from tqdm import tqdm |
7 | 10 | from typing import Union |
8 | 11 |
|
| 12 | +MULTIPART_CHUNK_SIZE = 10 * 1024 * 1024 # 10 MB |
| 13 | +MUST_USE_MULTIPART_THRESHOLD = 1024 * 1024 * 1024 # 1 GB |
| 14 | + |
9 | 15 |
|
10 | 16 | class Bucket(object): |
11 | 17 |
|
@@ -96,8 +102,150 @@ def get_file(self, name: str) -> DataproxyFile: |
96 | 102 | return file |
97 | 103 | raise DoesNotExist(f"Cannot find {name}.") |
98 | 104 |
|
| 105 | + def _get_filesize(self, filelike: Union[str, IOBase]) -> int: |
| 106 | + if isinstance(filelike, str): |
| 107 | + return os.path.getsize(filelike) |
| 108 | + pos = filelike.seek(0, 2) |
| 109 | + filelike.seek(0) |
| 110 | + return pos |
| 111 | + |
| 112 | + def _can_multipart_upload(self, filelike: Union[str, IOBase]) -> int: |
| 113 | + """Returns file size; raises InvalidParameter if not larger than MULTIPART_CHUNK_SIZE.""" |
| 114 | + size = self._get_filesize(filelike) |
| 115 | + if size <= MULTIPART_CHUNK_SIZE: |
| 116 | + raise InvalidParameter( |
| 117 | + f"multipart_upload requires file size > {MULTIPART_CHUNK_SIZE} bytes ({size} bytes given). Use upload() instead." |
| 118 | + ) |
| 119 | + return size |
| 120 | + |
| 121 | + @on_401_raise_unauthorized("Unauthorized") |
| 122 | + def multipart_upload(self, filelike: Union[str, IOBase], filename: str, **kwargs): |
| 123 | + """ |
| 124 | + Upload a file using multipart upload to the ebrains drive, supporting resumable uploads. |
| 125 | +
|
| 126 | + This method handles large file uploads by splitting them into chunks and uploading |
| 127 | + each chunk separately using presigned URLs. It supports resuming interrupted uploads |
| 128 | + via a local manifest file (`.multipart_manifest.json`), storing upload progress and |
| 129 | + ETags for completed parts. |
| 130 | +
|
| 131 | + Parameters |
| 132 | + ---------- |
| 133 | + filelike : str or IOBase |
| 134 | + Path to the local file as a string, or a file-like object (e.g., open file handle). |
| 135 | + If a string, the method opens and reads the file; if a file-like object, it must |
| 136 | + support `.seek()` and `.read()` operations. |
| 137 | + filename : str |
| 138 | + Destination filename in the remote storage. Leading slashes are stripped. |
| 139 | + **kwargs : dict, optional |
| 140 | + Additional keyword arguments (currently unused; included for extensibility). |
| 141 | +
|
| 142 | + Notes |
| 143 | + ----- |
| 144 | + - The manifest file is named `<filepath>.multipart_manifest.json` when a file path is provided. |
| 145 | + - The manifest stores: |
| 146 | + - `upload_id`: ID returned by the server for the multipart upload session. |
| 147 | + - `etag_maps`: Dictionary mapping part numbers to ETags of completed parts. |
| 148 | + - `next_offset`: Byte offset for the next chunk to upload (used for resuming). |
| 149 | + - In case of interruption, the method resumes from the last saved `next_offset`. |
| 150 | + - The manifest file is deleted upon successful completion of the upload. |
| 151 | +
|
| 152 | + Raises |
| 153 | + ------ |
| 154 | + UpstreamAPIException |
| 155 | + If the upload ID cannot be obtained from the server, or if a presigned URL |
| 156 | + is missing for a part, or if other required responses are malformed. |
| 157 | + """ |
| 158 | + sess = requests.Session() |
| 159 | + filename = filename.lstrip("/") |
| 160 | + self._can_multipart_upload(filelike) |
| 161 | + |
| 162 | + filepath = filelike if isinstance(filelike, str) else None |
| 163 | + manifest_path = f"{filepath}.multipart_manifest.json" if filepath else None |
| 164 | + |
| 165 | + # Load manifest for resume, or start fresh |
| 166 | + manifest = {} |
| 167 | + if manifest_path and os.path.exists(manifest_path): |
| 168 | + with open(manifest_path) as f: |
| 169 | + manifest = json.load(f) |
| 170 | + |
| 171 | + upload_id = manifest.get("upload_id") |
| 172 | + etag_maps: dict = manifest.get("etag_maps", {}) |
| 173 | + |
| 174 | + if not upload_id: |
| 175 | + resp = self.client.put(f"/v1/{self.target}/{self.dataproxy_entity_name}/{filename}/multipart") |
| 176 | + upload_id = resp.json().get("uploadId") |
| 177 | + if not upload_id: |
| 178 | + raise UpstreamAPIException("multipart_upload: failed to obtain uploadId.") |
| 179 | + manifest = {"upload_id": upload_id, "etag_maps": {}, "next_offset": 0} |
| 180 | + if manifest_path: |
| 181 | + with open(manifest_path, "w") as f: |
| 182 | + json.dump(manifest, f) |
| 183 | + |
| 184 | + next_offset = manifest.get("next_offset", 0) |
| 185 | + part_number = len(etag_maps) + 1 |
| 186 | + file_size = self._get_filesize(filelike) |
| 187 | + |
| 188 | + filehandle = open(filepath, "rb") if filepath else filelike |
| 189 | + try: |
| 190 | + filehandle.seek(next_offset) |
| 191 | + with tqdm( |
| 192 | + total=file_size, initial=next_offset, unit="B", unit_scale=True, unit_divisor=1024, desc=filename |
| 193 | + ) as progress: |
| 194 | + while True: |
| 195 | + chunk = filehandle.read(MULTIPART_CHUNK_SIZE) |
| 196 | + if not chunk: |
| 197 | + break |
| 198 | + resp = self.client.put( |
| 199 | + f"/v1/{self.target}/{self.dataproxy_entity_name}/{filename}/multipart/{upload_id}/{part_number}", |
| 200 | + params={"redirect": "false"}, |
| 201 | + ) |
| 202 | + part_url = resp.json().get("url") |
| 203 | + if not part_url: |
| 204 | + raise UpstreamAPIException(f"multipart_upload: no presigned URL for part {part_number}.") |
| 205 | + resp = sess.put(part_url, data=chunk) |
| 206 | + resp.raise_for_status() |
| 207 | + etag = resp.headers.get("etag", "").strip('"') |
| 208 | + etag_maps[str(part_number)] = etag |
| 209 | + next_offset += len(chunk) |
| 210 | + progress.update(len(chunk)) |
| 211 | + if manifest_path: |
| 212 | + manifest["etag_maps"] = etag_maps |
| 213 | + manifest["next_offset"] = next_offset |
| 214 | + with open(manifest_path, "w") as f: |
| 215 | + json.dump(manifest, f) |
| 216 | + part_number += 1 |
| 217 | + finally: |
| 218 | + if filepath: |
| 219 | + filehandle.close() |
| 220 | + |
| 221 | + resp = self.client.put( |
| 222 | + f"/v1/{self.target}/{self.dataproxy_entity_name}/{filename}/multipart/{upload_id}", |
| 223 | + params={"redirect": "false"}, |
| 224 | + json=etag_maps, |
| 225 | + ) |
| 226 | + |
| 227 | + if manifest_path and os.path.exists(manifest_path): |
| 228 | + os.remove(manifest_path) |
| 229 | + |
99 | 230 | @on_401_raise_unauthorized("Unauthorized") |
100 | 231 | def upload(self, filelike: Union[str, IOBase], filename: str, **kwargs): |
| 232 | + """ |
| 233 | + Upload a file to the bucket. If the file size exceeds `MUST_USE_MULTIPART_THRESHOLD` (1GB), |
| 234 | + this method automatically uses `multipart_upload` instead of a simple upload. |
| 235 | +
|
| 236 | + Parameters |
| 237 | + ---------- |
| 238 | + filelike : str or IOBase |
| 239 | + Path to the file or a file-like object (e.g., opened file handle). |
| 240 | + filename : str |
| 241 | + Destination filename in the bucket (leading slashes are stripped). |
| 242 | + **kwargs : dict, optional |
| 243 | + Additional keyword arguments passed to the underlying HTTP requests |
| 244 | + (e.g., headers, timeout, etc.). |
| 245 | + """ |
| 246 | + if self._get_filesize(filelike) > MUST_USE_MULTIPART_THRESHOLD: |
| 247 | + # use multipart upload to stay way below 5G gateway limit |
| 248 | + return self.multipart_upload(filelike, filename, **kwargs) |
101 | 249 | filename = filename.lstrip("/") |
102 | 250 | resp = self.client.put(f"/v1/{self.target}/{self.dataproxy_entity_name}/{filename}", **kwargs) |
103 | 251 | upload_url = resp.json().get("url") |
|
0 commit comments