Fix S3 signature generation
This commit is contained in:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user