diff --git a/classquiz/routers/storage.py b/classquiz/routers/storage.py index 7c64025..e20bd02 100644 --- a/classquiz/routers/storage.py +++ b/classquiz/routers/storage.py @@ -118,7 +118,9 @@ async def download_file_head(file_name: str) -> Response: @router.post("/") -async def upload_file(file: UploadFile = File(), user: User = Depends(get_current_user)) -> PublicStorageItem: +async def upload_file( + request: Request, file: UploadFile = File(), user: User = Depends(get_current_user) +) -> PublicStorageItem: if file.content_type not in ALLOWED_MIME_TYPES: raise HTTPException(status_code=422, detail="Unsupported") if user.storage_used > settings.free_storage_limit: @@ -135,7 +137,12 @@ async def upload_file(file: UploadFile = File(), user: User = Depends(get_curren deleted_at=None, alt_text=None, ) - await storage.upload(file_name=file_id.hex, file_data=file.file, mime_type=file.content_type) + await storage.upload( + file_name=file_id.hex, + file_data=file.file, + mime_type=file.content_type, + size=request.headers.get("Content-Length"), + ) await file_obj.save() await arq.enqueue_job("calculate_hash", file_id.hex) return PublicStorageItem.from_db_model(file_obj) diff --git a/classquiz/storage/__init__.py b/classquiz/storage/__init__.py index e90e9c7..90b277c 100644 --- a/classquiz/storage/__init__.py +++ b/classquiz/storage/__init__.py @@ -46,8 +46,10 @@ class Storage: def download(self, file_name: str) -> Generator | None: return self.instance.download(file_name) - async def upload(self, file_name: str, file_data: BinaryIO, mime_type: str | None = None) -> None: - return await self.instance.upload(file=file_data, file_name=file_name, mime_type=mime_type) + async def upload( + self, file_name: str, file_data: BinaryIO, mime_type: str | None = None, size: int | None = None + ) -> None: + return await self.instance.upload(file=file_data, file_name=file_name, mime_type=mime_type, size=size) async def delete(self, file_names: [str]) -> None: return await self.instance.delete(file_names=file_names) diff --git a/classquiz/storage/local_storage.py b/classquiz/storage/local_storage.py index 3a6f1d4..1404e11 100644 --- a/classquiz/storage/local_storage.py +++ b/classquiz/storage/local_storage.py @@ -29,7 +29,13 @@ class LocalStorage: yield None # skipcq: PYL-W0613 - async def upload(self, file_name: str, file: BinaryIO, mime_type: str | None = None) -> None: + async def upload( + self, + file_name: str, + file: BinaryIO, + size: int | None, + mime_type: str | None = None, + ) -> None: with open(file=os.path.join(self.base_path, file_name), mode="wb") as f: copyfileobj(file, f) diff --git a/classquiz/storage/s3_storage.py b/classquiz/storage/s3_storage.py index 913d80e..af2832b 100644 --- a/classquiz/storage/s3_storage.py +++ b/classquiz/storage/s3_storage.py @@ -5,7 +5,6 @@ import hashlib import hmac -import sys from datetime import datetime, timedelta from typing import Tuple, BinaryIO, Generator @@ -109,9 +108,9 @@ class S3Storage: return headers, request_url # skipcq: PYL-W0613 - async def upload(self, file: BinaryIO, file_name: str, mime_type: str | None = "application/octet-stream") -> None: + async def upload(self, file: BinaryIO, file_name: str, size: int | None, mime_type: str | None = None) -> None: headers, url = self._generate_aws_signature_v4(method="PUT", path=f"/{file_name}") - headers["Content-Length"] = sys.getsizeof(file) + headers["Content-Length"] = size async with ClientSession() as session, session.put(url, headers=headers, data=file) as resp: if resp.status == 200: return None