Fix S3 signature generation

This commit is contained in:
Mawoka
2025-11-01 13:53:36 +01:00
parent 6faf321c2a
commit 048f5d3e7b
+80 -56
View File
@@ -11,7 +11,11 @@ from typing import Tuple, BinaryIO, Generator
from aiohttp import ClientSession
import minio
from pydantic import BaseModel
from classquiz.storage.errors import DeletionFailedError, SavingFailedError, DownloadingFailedError
from classquiz.storage.errors import (
DeletionFailedError,
SavingFailedError,
DownloadingFailedError,
)
class S3Storage:
@@ -19,7 +23,14 @@ class S3Storage:
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"):
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
@@ -35,80 +46,82 @@ class S3Storage:
path = f"/{self.bucket_name}{path}"
service = "s3"
# Create a timestamp for the request
# --- Timestamp ---
t = datetime.utcnow()
amz_date = t.strftime("%Y%m%dT%H%M%SZ")
datestamp = t.strftime("%Y%m%d")
# Create a canonical request
# --- Canonical request parts ---
canonical_uri = path
canonical_querystring = ""
if expiry is not None:
canonical_querystring = f"Expires={expiry}"
canonical_headers = "host:" + self.host + "\n" + "x-amz-date:" + amz_date + "\n"
signed_headers = "host;x-amz-date"
payload_hash = hashlib.sha256("".encode("utf-8")).hexdigest()
# For S3, the payload hash must be included and signed
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 = (
method
+ "\n"
+ canonical_uri
+ "\n"
+ canonical_querystring
+ "\n"
+ canonical_headers
+ "\n"
+ signed_headers
+ "\n"
+ payload_hash
f"{method}\n"
f"{canonical_uri}\n"
f"{canonical_querystring}\n"
f"{canonical_headers}\n"
f"{signed_headers}\n"
f"{payload_hash}"
)
# Create a string to sign
# --- String to sign ---
algorithm = "AWS4-HMAC-SHA256"
credential_scope = datestamp + "/" + self.region + "/" + service + "/" + "aws4_request"
credential_scope = f"{datestamp}/{self.region}/{service}/aws4_request"
string_to_sign = (
algorithm
+ "\n"
+ amz_date
+ "\n"
+ credential_scope
+ "\n"
+ hashlib.sha256(canonical_request.encode("utf-8")).hexdigest()
f"{algorithm}\n"
f"{amz_date}\n"
f"{credential_scope}\n"
f"{hashlib.sha256(canonical_request.encode('utf-8')).hexdigest()}"
)
# Create a signing key
k_date = hmac.new(
("AWS4" + self.secret_key).encode("utf-8"), datestamp.encode("utf-8"), hashlib.sha256
).digest()
k_region = hmac.new(k_date, self.region.encode("utf-8"), hashlib.sha256).digest()
k_service = hmac.new(k_region, service.encode("utf-8"), hashlib.sha256).digest()
k_signing = hmac.new(k_service, b"aws4_request", hashlib.sha256).digest()
# --- Derive signing key ---
def sign(key, msg):
return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest()
# Calculate the signature
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()
# Add the signature to the request as an Authorization header
# --- Authorization header ---
authorization_header = (
algorithm
+ " "
+ "Credential="
+ self.access_key
+ "/"
+ credential_scope
+ ", "
+ "SignedHeaders="
+ signed_headers
+ ", "
+ "Signature="
+ signature
f"{algorithm} "
f"Credential={self.access_key}/{credential_scope}, "
f"SignedHeaders={signed_headers}, "
f"Signature={signature}"
)
# Send the request with the authorization header
headers = {"x-amz-date": amz_date, "Authorization": authorization_header}
request_url = self.base_url + path + "?" + canonical_querystring
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:
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}")
file_size = 0
while True:
@@ -119,7 +132,10 @@ class S3Storage:
break
headers["Content-Length"] = str(file_size)
async with ClientSession() as session, session.put(url, headers=headers, data=file) as resp:
async with (
ClientSession() as session,
session.put(url, headers=headers, data=file) as resp,
):
if resp.status == 200:
return None
else:
@@ -129,7 +145,10 @@ class S3Storage:
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:
async with (
ClientSession() as session,
session.delete(url, headers=headers) as resp,
):
if resp.status == 204:
return None
else:
@@ -137,7 +156,9 @@ class S3Storage:
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)
object_name=file_name,
bucket_name=self.bucket_name,
expires=timedelta(seconds=expire),
)
def size(self, file_name: str) -> int | None:
@@ -150,7 +171,10 @@ class S3Storage:
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:
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