"""Utilities to work with S3.""" import asyncio import io from tempfile import NamedTemporaryFile from typing import Any, Dict, List from urllib.parse import quote from openpyxl import Workbook from types_aiobotocore_s3.type_defs import ( GetObjectOutputTypeDef, HeadObjectOutputTypeDef, ) from product_staging import config from product_staging.api import datasources from product_staging.api.schemas.assets_upload import ETagPart async def get_multipart_upload_presigned_urls( parts: int, s3_filename: str, upload_token: str ) -> List[Dict[str, Any]]: """Return presigned URLs for a number of parts, a s3_filename and a token associated with the upload.""" s3_client = datasources.get_s3_client() async def get_part_presigned_url(part_number: int) -> Dict[str, Any]: return { "part_number": part_number, "url": await s3_client.generate_presigned_url( ClientMethod="upload_part", Params={ "Bucket": config.OWS_PRODUCT_STAGING_S3_BUCKET, "Key": s3_filename, "UploadId": upload_token, "PartNumber": part_number, }, ), } return await asyncio.gather( *[get_part_presigned_url(p) for p in range(1, parts + 1)] ) async def get_file(s3_filename: str) -> GetObjectOutputTypeDef: return await datasources.get_s3_client().get_object( Bucket=config.OWS_PRODUCT_STAGING_S3_BUCKET, Key=s3_filename ) async def head_file(s3_filename: str) -> HeadObjectOutputTypeDef: return await datasources.get_s3_client().head_object( Bucket=config.OWS_PRODUCT_STAGING_S3_BUCKET, Key=s3_filename ) def get_file_metadata(file_obj: GetObjectOutputTypeDef) -> Dict[str, str]: return file_obj["Metadata"] async def get_file_stream(s3_filename: str) -> io.BytesIO: return io.BytesIO(await (await get_file(s3_filename))["Body"].read()) async def write_file_stream( s3_filename: str, stream: io.BytesIO, metadata: Dict[str, str] ) -> None: await datasources.get_s3_client().put_object( Bucket=config.OWS_PRODUCT_STAGING_S3_BUCKET, Key=s3_filename, Body=stream.getbuffer().tobytes(), Metadata=metadata, ) async def delete_file(s3_filename: str) -> Any: return await datasources.get_s3_client().delete_object( Bucket=config.OWS_PRODUCT_STAGING_S3_BUCKET, Key=s3_filename ) def _content_disposition(filename: str) -> str: """Build a Content-Disposition header value safe for S3 (ISO-8859-1). Always uses RFC 5987 encoding (filename*=UTF-8''...) to avoid any character-set issues with S3's ISO-8859-1 restriction. """ # Fallback for backwards compatibility with records that have no filename filename = filename or "download" encoded = quote(filename, safe="") return f"attachment; filename=\"{encoded}\"; filename*=UTF-8''{encoded}" async def upload_workbook_and_generate_download_link( workbook: Workbook, dest_filename: str, folder_name: str ): """ Upload spreadsheet file to s3 and return download link """ s3_client = datasources.get_s3_client() with NamedTemporaryFile(): filename = f"/tmp/{dest_filename}" workbook.save(filename) await s3_client.upload_file( Bucket=config.OWS_PRODUCT_STAGING_S3_BUCKET, Filename=filename, Key=f"{folder_name}/{dest_filename}", ) download_link = await s3_client.generate_presigned_url( "get_object", Params={ "Bucket": config.OWS_PRODUCT_STAGING_S3_BUCKET, "Key": f"{folder_name}/{dest_filename}", "ResponseContentDisposition": _content_disposition(dest_filename), }, ExpiresIn=86400, ) return download_link async def generate_download_link(key: str, filename: str | None = None): s3_client = datasources.get_s3_client() download_link = await s3_client.generate_presigned_url( "get_object", Params={ "Bucket": config.OWS_PRODUCT_STAGING_S3_BUCKET, "Key": key, "ResponseContentDisposition": _content_disposition(filename or key), }, ExpiresIn=86400, ) return download_link async def complete_multipart_upload( s3_filename: str, upload_token: str, parts: list[ETagPart] ): s3_client = datasources.get_s3_client() max_retries = 3 for attempt in range(max_retries + 1): try: return await s3_client.complete_multipart_upload( Bucket=config.OWS_PRODUCT_STAGING_S3_BUCKET, Key=s3_filename, UploadId=upload_token, MultipartUpload={ "Parts": [ {"PartNumber": part.part_number, "ETag": part.etag} for part in parts ] }, ) except Exception as e: last_exception = e if attempt == max_retries: break await asyncio.sleep(1 + attempt) if last_exception is not None: raise last_exception