Files
classquiz-ai/classquiz/storage/s3_storage.py
T
2025-11-01 14:06:26 +01:00

200 lines
6.3 KiB
Python

# SPDX-FileCopyrightText: 2023 Marlon W (Mawoka)
#
# SPDX-License-Identifier: MPL-2.0
import hashlib
import hmac
from datetime import datetime, timedelta
from typing import Tuple, BinaryIO, Generator
from aiohttp import ClientSession
import minio
from pydantic import BaseModel
from classquiz.storage.errors import (
DeletionFailedError,
SavingFailedError,
DownloadingFailedError,
)
class S3Storage:
class _HeaderAndParams(BaseModel):
params: dict[str, str]
headers: dict[str, str]
def __init__(
self,
base_url: str,
access_key: str,
secret_key: str,
bucket_name: str,
region: str = "us-east-1",
):
self.base_url = base_url
self.access_key = access_key
self.secret_key = secret_key
self.bucket_name = bucket_name
self.region = region
self.DATE_FORMAT = "%a, %d %b %Y %H:%M:%S GMT"
self.host = base_url.replace("http://", "").replace("https://", "")
self.client = minio.Minio(self.host, access_key=access_key, secret_key=secret_key)
if not self.client.bucket_exists(self.bucket_name):
self.client.make_bucket(self.bucket_name)
def _generate_aws_signature_v4(
self,
method: str,
path: str,
expiry: int | None = None,
payload_hash: str | None = None,
) -> Tuple[dict, str]:
path = f"/{self.bucket_name}{path}"
service = "s3"
# --- Timestamp ---
t = datetime.utcnow()
amz_date = t.strftime("%Y%m%dT%H%M%SZ")
datestamp = t.strftime("%Y%m%d")
# --- Canonical request parts ---
canonical_uri = path
canonical_querystring = ""
if expiry is not None:
canonical_querystring = f"Expires={expiry}"
# For S3, the payload hash must be included and signed
if payload_hash is None:
payload_hash = hashlib.sha256(b"").hexdigest()
canonical_headers = f"host:{self.host}\n" f"x-amz-content-sha256:{payload_hash}\n" f"x-amz-date:{amz_date}\n"
signed_headers = "host;x-amz-content-sha256;x-amz-date"
canonical_request = (
f"{method}\n"
f"{canonical_uri}\n"
f"{canonical_querystring}\n"
f"{canonical_headers}\n"
f"{signed_headers}\n"
f"{payload_hash}"
)
# --- String to sign ---
algorithm = "AWS4-HMAC-SHA256"
credential_scope = f"{datestamp}/{self.region}/{service}/aws4_request"
string_to_sign = (
f"{algorithm}\n"
f"{amz_date}\n"
f"{credential_scope}\n"
f"{hashlib.sha256(canonical_request.encode('utf-8')).hexdigest()}"
)
# --- Derive signing key ---
def sign(key, msg):
return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest()
k_date = sign(("AWS4" + self.secret_key).encode("utf-8"), datestamp)
k_region = sign(k_date, self.region)
k_service = sign(k_region, service)
k_signing = sign(k_service, "aws4_request")
# --- Signature ---
signature = hmac.new(k_signing, string_to_sign.encode("utf-8"), hashlib.sha256).hexdigest()
# --- Authorization header ---
authorization_header = (
f"{algorithm} "
f"Credential={self.access_key}/{credential_scope}, "
f"SignedHeaders={signed_headers}, "
f"Signature={signature}"
)
headers = {
"x-amz-date": amz_date,
"x-amz-content-sha256": payload_hash,
"Authorization": authorization_header,
}
request_url = self.base_url + path
if canonical_querystring:
request_url += f"?{canonical_querystring}"
return headers, request_url
# skipcq: PYL-W0613
async def upload(
self,
file: BinaryIO,
file_name: str,
size: int | None,
mime_type: str | None = None,
) -> None:
file.seek(0)
payload = file.read()
payload_hash = hashlib.sha256(payload).hexdigest()
file.seek(0)
headers, url = self._generate_aws_signature_v4(method="PUT", path=f"/{file_name}", payload_hash=payload_hash)
file_size = 0
while True:
chunk = file.read(1024)
file_size += len(chunk)
if not chunk:
file.seek(0)
break
headers["Content-Length"] = str(file_size)
async with (
ClientSession() as session,
session.put(url, headers=headers, data=file) as resp,
):
if resp.status == 200:
return None
else:
print(await resp.text())
raise SavingFailedError
async def delete(self, file_names: list[str]) -> None:
for file in file_names:
headers, url = self._generate_aws_signature_v4(method="DELETE", path=f"/{file}")
async with (
ClientSession() as session,
session.delete(url, headers=headers) as resp,
):
if resp.status == 204:
return None
else:
raise DeletionFailedError
def get_url(self, expire: int, file_name: str) -> str:
return self.client.presigned_get_object(
object_name=file_name,
bucket_name=self.bucket_name,
expires=timedelta(seconds=expire),
)
def size(self, file_name: str) -> int | None:
try:
res = self.client.stat_object(bucket_name=self.bucket_name, object_name=file_name)
except minio.error.S3Error:
return None
return res.size
async def download(self, file_name: str) -> Generator:
headers, url = self._generate_aws_signature_v4(method="GET", path=f"/{file_name}")
async with (
ClientSession() as session,
session.get(url, headers=headers) as resp,
):
if resp.status == 200:
async for i in resp.content.iter_chunked(1024):
yield i
elif resp.status == 404:
yield None
else:
raise DownloadingFailedError
# client = httpx.AsyncClient()
# async with client.stream("GET", url, headers=headers) as resp:
# if resp.status == 200:
# yield resp.aiter_bytes()