117 lines
3.5 KiB
Python
117 lines
3.5 KiB
Python
import base64
|
|
import hashlib
|
|
import hmac
|
|
from time import time
|
|
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
|
|
|
|
from .config import get_settings
|
|
|
|
|
|
SIGNED_QUERY_KEYS = {"OSSAccessKeyId", "Expires", "Signature"}
|
|
MEDIA_URL_EXPIRES_SECONDS = 3600
|
|
|
|
|
|
def normalized_oss_host(endpoint: str, bucket: str) -> tuple[str, str]:
|
|
raw = endpoint.strip().rstrip("/")
|
|
if "://" not in raw:
|
|
raw = f"https://{raw}"
|
|
parsed = urlsplit(raw)
|
|
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
|
raise ValueError("invalid OSS endpoint")
|
|
host = parsed.netloc
|
|
if not host.lower().startswith(f"{bucket.lower()}."):
|
|
host = f"{bucket}.{host}"
|
|
return parsed.scheme, host
|
|
|
|
|
|
def normalized_public_base_url(base_url: str) -> tuple[str, str]:
|
|
raw = base_url.strip().rstrip("/")
|
|
if "://" not in raw:
|
|
raw = f"https://{raw}"
|
|
parsed = urlsplit(raw)
|
|
if (
|
|
parsed.scheme not in {"http", "https"}
|
|
or not parsed.netloc
|
|
or parsed.username
|
|
or parsed.password
|
|
or parsed.path not in {"", "/"}
|
|
or parsed.query
|
|
or parsed.fragment
|
|
):
|
|
raise ValueError("invalid OSS public base URL")
|
|
return parsed.scheme, parsed.netloc
|
|
|
|
|
|
def sign_oss_get_url(
|
|
url: str,
|
|
*,
|
|
access_key_id: str,
|
|
access_key_secret: str,
|
|
endpoint: str,
|
|
bucket: str,
|
|
expires_at: int,
|
|
public_base_url: str | None = None,
|
|
) -> str:
|
|
source_scheme, source_host = normalized_oss_host(endpoint, bucket)
|
|
public_scheme, public_host = (
|
|
normalized_public_base_url(public_base_url)
|
|
if public_base_url and public_base_url.strip()
|
|
else (source_scheme, source_host)
|
|
)
|
|
parsed = urlsplit(url)
|
|
configured_hosts = {source_host.lower(), public_host.lower()}
|
|
if parsed.scheme not in {"http", "https"} or parsed.netloc.lower() not in configured_hosts:
|
|
return url
|
|
|
|
path = parsed.path or "/"
|
|
canonical_resource = f"/{bucket}{path}"
|
|
string_to_sign = f"GET\n\n\n{expires_at}\n{canonical_resource}"
|
|
signature = base64.b64encode(
|
|
hmac.new(
|
|
access_key_secret.strip().encode("utf-8"),
|
|
string_to_sign.encode("utf-8"),
|
|
hashlib.sha1,
|
|
).digest()
|
|
).decode("ascii")
|
|
query = [
|
|
(key, value)
|
|
for key, value in parse_qsl(parsed.query, keep_blank_values=True)
|
|
if key not in SIGNED_QUERY_KEYS
|
|
]
|
|
query.extend(
|
|
[
|
|
("OSSAccessKeyId", access_key_id.strip()),
|
|
("Expires", str(expires_at)),
|
|
("Signature", signature),
|
|
]
|
|
)
|
|
return urlunsplit((public_scheme, public_host, path, urlencode(query), ""))
|
|
|
|
|
|
def resolve_media_url(url: str | None) -> str | None:
|
|
if not url or not isinstance(url, str):
|
|
return url
|
|
|
|
settings = get_settings()
|
|
values = (
|
|
settings.oss_access_key_id,
|
|
settings.oss_access_key_secret,
|
|
settings.oss_endpoint,
|
|
settings.oss_bucket_name,
|
|
)
|
|
if not all(value and value.strip() for value in values):
|
|
return url
|
|
|
|
try:
|
|
return sign_oss_get_url(
|
|
url,
|
|
access_key_id=settings.oss_access_key_id,
|
|
access_key_secret=settings.oss_access_key_secret,
|
|
endpoint=settings.oss_endpoint,
|
|
bucket=settings.oss_bucket_name,
|
|
expires_at=int(time()) + MEDIA_URL_EXPIRES_SECONDS,
|
|
public_base_url=getattr(settings, "oss_public_base_url", None),
|
|
)
|
|
except (TypeError, ValueError):
|
|
return url
|