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 from aiohttp import ClientSession
import minio import minio
from pydantic import BaseModel from pydantic import BaseModel
from classquiz.storage.errors import DeletionFailedError, SavingFailedError, DownloadingFailedError from classquiz.storage.errors import (
DeletionFailedError,
SavingFailedError,
DownloadingFailedError,
)
class S3Storage: class S3Storage:
@@ -19,7 +23,14 @@ class S3Storage:
params: dict[str, str] params: dict[str, str]
headers: 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.base_url = base_url
self.access_key = access_key self.access_key = access_key
self.secret_key = secret_key self.secret_key = secret_key
@@ -35,80 +46,82 @@ class S3Storage:
path = f"/{self.bucket_name}{path}" path = f"/{self.bucket_name}{path}"
service = "s3" service = "s3"
# Create a timestamp for the request # --- Timestamp ---
t = datetime.utcnow() t = datetime.utcnow()
amz_date = t.strftime("%Y%m%dT%H%M%SZ") amz_date = t.strftime("%Y%m%dT%H%M%SZ")
datestamp = t.strftime("%Y%m%d") datestamp = t.strftime("%Y%m%d")
# Create a canonical request # --- Canonical request parts ---
canonical_uri = path canonical_uri = path
canonical_querystring = "" canonical_querystring = ""
if expiry is not None: if expiry is not None:
canonical_querystring = f"Expires={expiry}" canonical_querystring = f"Expires={expiry}"
canonical_headers = "host:" + self.host + "\n" + "x-amz-date:" + amz_date + "\n"
signed_headers = "host;x-amz-date" # For S3, the payload hash must be included and signed
payload_hash = hashlib.sha256("".encode("utf-8")).hexdigest() 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 = ( canonical_request = (
method f"{method}\n"
+ "\n" f"{canonical_uri}\n"
+ canonical_uri f"{canonical_querystring}\n"
+ "\n" f"{canonical_headers}\n"
+ canonical_querystring f"{signed_headers}\n"
+ "\n" f"{payload_hash}"
+ canonical_headers
+ "\n"
+ signed_headers
+ "\n"
+ payload_hash
) )
# Create a string to sign # --- String to sign ---
algorithm = "AWS4-HMAC-SHA256" algorithm = "AWS4-HMAC-SHA256"
credential_scope = datestamp + "/" + self.region + "/" + service + "/" + "aws4_request" credential_scope = f"{datestamp}/{self.region}/{service}/aws4_request"
string_to_sign = ( string_to_sign = (
algorithm f"{algorithm}\n"
+ "\n" f"{amz_date}\n"
+ amz_date f"{credential_scope}\n"
+ "\n" f"{hashlib.sha256(canonical_request.encode('utf-8')).hexdigest()}"
+ credential_scope
+ "\n"
+ hashlib.sha256(canonical_request.encode("utf-8")).hexdigest()
) )
# Create a signing key # --- Derive signing key ---
k_date = hmac.new( def sign(key, msg):
("AWS4" + self.secret_key).encode("utf-8"), datestamp.encode("utf-8"), hashlib.sha256 return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest()
).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()
# 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() 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 = ( authorization_header = (
algorithm f"{algorithm} "
+ " " f"Credential={self.access_key}/{credential_scope}, "
+ "Credential=" f"SignedHeaders={signed_headers}, "
+ self.access_key f"Signature={signature}"
+ "/"
+ credential_scope
+ ", "
+ "SignedHeaders="
+ signed_headers
+ ", "
+ "Signature="
+ signature
) )
# Send the request with the authorization header
headers = {"x-amz-date": amz_date, "Authorization": authorization_header} headers = {
request_url = self.base_url + path + "?" + canonical_querystring "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 return headers, request_url
# skipcq: PYL-W0613 # 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}") headers, url = self._generate_aws_signature_v4(method="PUT", path=f"/{file_name}")
file_size = 0 file_size = 0
while True: while True:
@@ -119,7 +132,10 @@ class S3Storage:
break break
headers["Content-Length"] = str(file_size) 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: if resp.status == 200:
return None return None
else: else:
@@ -129,7 +145,10 @@ class S3Storage:
async def delete(self, file_names: list[str]) -> None: async def delete(self, file_names: list[str]) -> None:
for file in file_names: for file in file_names:
headers, url = self._generate_aws_signature_v4(method="DELETE", path=f"/{file}") 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: if resp.status == 204:
return None return None
else: else:
@@ -137,7 +156,9 @@ class S3Storage:
def get_url(self, expire: int, file_name: str) -> str: def get_url(self, expire: int, file_name: str) -> str:
return self.client.presigned_get_object( 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: def size(self, file_name: str) -> int | None:
@@ -150,7 +171,10 @@ class S3Storage:
async def download(self, file_name: str) -> Generator: async def download(self, file_name: str) -> Generator:
headers, url = self._generate_aws_signature_v4(method="GET", path=f"/{file_name}") 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: if resp.status == 200:
async for i in resp.content.iter_chunked(1024): async for i in resp.content.iter_chunked(1024):
yield i yield i