diff --git a/python-client/wmill/wmill/client.py b/python-client/wmill/wmill/client.py index 3ac04a0318..c212af21b6 100644 --- a/python-client/wmill/wmill/client.py +++ b/python-client/wmill/wmill/client.py @@ -3,6 +3,7 @@ from __future__ import annotations import atexit import datetime as dt import functools +from io import BufferedReader, BytesIO import logging import os import random @@ -395,7 +396,7 @@ class Windmill: part_response = self.post( f"/w/{self.workspace}/job_helpers/multipart_download_s3_file", json={ - "file_key": s3object.s3, + "file_key": s3object["s3"], "part_number": part_number, "file_size": file_total_size, "s3_resource_path": s3_resource_path, @@ -413,7 +414,7 @@ class Windmill: def write_s3_file( self, s3object: S3Object | None, - file_content: bytes, + file_content: BufferedReader | bytes, file_expiration: dt.datetime | None, s3_resource_path: str | None, ) -> S3Object: @@ -424,26 +425,56 @@ class Windmill: from wmill import S3Object s3_obj = S3Object(s3="/path/to/my_file.txt") + + # for an in memory bytes array: file_content = b'Hello Windmill!' client.write_s3_file(s3_obj, file_content) + + # for a file: + with open("my_file.txt", "rb") as my_file: + client.write_s3_file(s3_obj, my_file) ''' """ - try: - result = self.post( - f"/w/{self.workspace}/job_helpers/multipart_upload_s3_file", - json={ - "file_key": s3object.s3 if s3object is not None else None, - "part_content": file_content, - "parts": [], - "is_final": True, - "cancel_upload": False, - "s3_resource_path": s3_resource_path if s3_resource_path != "" else None, - "file_expiration": file_expiration.isoformat() if file_expiration else None, - }, - ).json() - except Exception as e: - raise Exception("Could not write file to S3") from e - return S3Object(s3=result["file_key"]) + content_reader: BufferedReader | BytesIO + if isinstance(file_content, BufferedReader): + content_reader = file_content + elif isinstance(file_content, bytes): + content_reader = BytesIO(file_content) + else: + raise Exception("Type of file_content not supported") + + file_key = s3object["s3"] if s3object is not None else None + parts = [] + upload_id = None + chunk = content_reader.read(5 * 1024 * 1024) + if len(chunk) == 0: + raise Exception("File content is empty, nothing to upload") + while True: + chunk_2 = content_reader.read(5 * 1024 * 1024) + reader_done = len(chunk_2) == 0 + try: + response = self.post( + f"/w/{self.workspace}/job_helpers/multipart_upload_s3_file", + json={ + "file_key": file_key, + "part_content": [b for b in chunk], + "upload_id": upload_id, + "parts": parts, + "is_final": reader_done, + "cancel_upload": False, + "s3_resource_path": s3_resource_path, + "file_expiration": file_expiration.isoformat() if file_expiration else None, + }, + ).json() + except Exception as e: + raise Exception("Could not write file to S3") from e + parts = response["parts"] + upload_id = response["upload_id"] + file_key = response["file_key"] + if response["is_done"]: + break + chunk = chunk_2 + return S3Object(s3=file_key) def __boto3_connection_settings(self, s3_resource) -> Boto3ConnectionSettings: endpoint_url_prefix = "https://" if s3_resource["useSSL"] else "http://" @@ -706,8 +737,8 @@ def load_s3_file(s3object: S3Object, s3_resource_path: str = "") -> bytes: @init_global_client def write_s3_file( s3object: S3Object | None, - file_content: bytes, - file_expiration: dt.datetime | None, + file_content: BufferedReader | bytes, + file_expiration: dt.datetime | None = None, s3_resource_path: str = "", ) -> S3Object: """