200 lines
6.3 KiB
Python
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()
|