mirror of
https://github.com/goauthentik/authentik.git
synced 2026-08-30 18:51:39 -07:00
* core: cache S3 file storage clients Agent-thread: https://sdko.org/internal/thr/ak/019ed837-6a7a-7492-bc76-93d55a130a27 A7k-product: product A7k-product-repo: 1 Co-authored-by: Agent <gptagent@svc.sdko.net> * core: limit S3 client cache refresh Agent-thread: https://koala.sdko.net/th?h=co&d=a7k&t=019edffd-14cc-7b70-abce-495e4af0a0b0 Signed-off-by: Dominic Roy <dominic@goauthentik.io> Co-authored-by: Agent <gptagent@svc.sdko.net> * core: trim S3 client cache tests Agent-thread: https://koala.sdko.net/th?h=co&d=a7k&t=019edffd-14cc-7b70-abce-495e4af0a0b0 Signed-off-by: Dominic Roy <dominic@goauthentik.io> Co-authored-by: Agent <gptagent@svc.sdko.net> * fix * Apply suggestion from @dominic-r Signed-off-by: Dominic R <dominic@goauthentik.io> * slim tests * fix * admin: refresh cached S3 clients on config changes Agent-thread: https://koala.sdko.net/th?h=co&d=a7k&t=019f04bc-8623-70d0-8db8-3888430aa6c8 Signed-off-by: Dominic Roy <dominic@goauthentik.io> Co-authored-by: Agent <gptagent@svc.sdko.net> --------- Signed-off-by: Dominic Roy <dominic@goauthentik.io> Signed-off-by: Dominic R <dominic@goauthentik.io> Co-authored-by: Agent <gptagent@svc.sdko.net>
285 lines
11 KiB
Python
285 lines
11 KiB
Python
from collections.abc import Generator, Iterator
|
|
from contextlib import contextmanager
|
|
from tempfile import SpooledTemporaryFile
|
|
from typing import Any, TypeVar
|
|
from urllib.parse import urlsplit, urlunsplit
|
|
|
|
import boto3
|
|
from botocore.config import Config
|
|
from botocore.exceptions import ClientError
|
|
from django.db import connection
|
|
from django.http.request import HttpRequest
|
|
|
|
from authentik.admin.files.backends.base import ManageableBackend, get_content_type
|
|
from authentik.admin.files.usage import FileUsage
|
|
from authentik.lib.config import CONFIG
|
|
from authentik.lib.utils.time import timedelta_from_string
|
|
|
|
_ConfigValue = TypeVar("_ConfigValue")
|
|
|
|
|
|
class S3Backend(ManageableBackend):
|
|
"""S3-compatible object storage backend.
|
|
|
|
Stores files in s3-compatible storage:
|
|
- Key prefix: {usage}/{schema}/{filename}
|
|
- Supports full file management (upload, delete, list)
|
|
- Generates presigned URLs for file access
|
|
- Used when storage.backend=s3
|
|
"""
|
|
|
|
allowed_usages = list(FileUsage) # All usages
|
|
name = "s3"
|
|
|
|
def __init__(self, *args, **kwargs):
|
|
super().__init__(*args, **kwargs)
|
|
self._config = {}
|
|
self._session = None
|
|
self._client = None
|
|
|
|
def _remember_config(self, key: str, refreshed: _ConfigValue) -> tuple[_ConfigValue, bool]:
|
|
unset = object()
|
|
current = self._config.get(key, unset)
|
|
if current is unset:
|
|
current = refreshed
|
|
self._config[key] = refreshed
|
|
return refreshed, current != refreshed
|
|
|
|
def _get_config(self, key: str, default: Any) -> tuple[Any, bool]:
|
|
refreshed = CONFIG.refresh(
|
|
f"storage.{self.usage.value}.{self.name}.{key}",
|
|
CONFIG.refresh(f"storage.{self.name}.{key}", default),
|
|
)
|
|
return self._remember_config(key, refreshed)
|
|
|
|
def _get_bool_config(self, key: str, default: bool) -> tuple[bool, bool]:
|
|
refreshed = CONFIG.get_bool(
|
|
f"storage.{self.usage.value}.{self.name}.{key}",
|
|
CONFIG.get_bool(f"storage.{self.name}.{key}", default),
|
|
)
|
|
return self._remember_config(key, refreshed)
|
|
|
|
@property
|
|
def base_path(self) -> str:
|
|
"""S3 key prefix: {usage}/{schema}/"""
|
|
return f"{self.usage.value}/{connection.schema_name}"
|
|
|
|
@property
|
|
def bucket_name(self) -> str:
|
|
return CONFIG.get(
|
|
f"storage.{self.usage.value}.{self.name}.bucket_name",
|
|
CONFIG.get(f"storage.{self.name}.bucket_name"),
|
|
)
|
|
|
|
@property
|
|
def session(self) -> boto3.Session:
|
|
"""Create boto3 session with configured credentials."""
|
|
session_profile, session_profile_r = self._get_config("session_profile", None)
|
|
if session_profile is not None:
|
|
if session_profile_r or self._session is None:
|
|
self._session = boto3.Session(profile_name=session_profile)
|
|
self._client = None
|
|
return self._session
|
|
else:
|
|
return self._session
|
|
else:
|
|
access_key, access_key_r = self._get_config("access_key", None)
|
|
secret_key, secret_key_r = self._get_config("secret_key", None)
|
|
session_token, session_token_r = self._get_config("session_token", None)
|
|
if access_key_r or secret_key_r or session_token_r or self._session is None:
|
|
self._session = boto3.Session(
|
|
aws_access_key_id=access_key,
|
|
aws_secret_access_key=secret_key,
|
|
aws_session_token=session_token,
|
|
)
|
|
self._client = None
|
|
return self._session
|
|
else:
|
|
return self._session
|
|
|
|
@property
|
|
def client(self):
|
|
"""Create S3 client with configured endpoint and region."""
|
|
endpoint_url, endpoint_url_r = self._get_config("endpoint", None)
|
|
session = self.session
|
|
use_ssl, use_ssl_r = self._get_bool_config("use_ssl", True)
|
|
region_name, region_name_r = self._get_config("region", None)
|
|
addressing_style, addressing_style_r = self._get_config("addressing_style", "auto")
|
|
signature_version, signature_version_r = self._get_config("signature_version", "s3v4")
|
|
|
|
if self._client is not None and not any(
|
|
(
|
|
endpoint_url_r,
|
|
use_ssl_r,
|
|
region_name_r,
|
|
addressing_style_r,
|
|
signature_version_r,
|
|
)
|
|
):
|
|
return self._client
|
|
# Keep signature_version pass-through and let boto3/botocore handle it.
|
|
# In boto3's S3 configuration docs, `s3v4` (default) and deprecated `s3`
|
|
# are the documented values:
|
|
# https://github.com/boto/boto3/blob/791a3e8f36d83664a47b4281a0586b3546cef3ec/docs/source/guide/configuration.rst?plain=1#L398-L407
|
|
# Botocore also supports additional signer names, so we intentionally do
|
|
# not enforce a restricted allowlist here.
|
|
|
|
self._client = session.client(
|
|
"s3",
|
|
endpoint_url=endpoint_url,
|
|
use_ssl=use_ssl,
|
|
region_name=region_name,
|
|
config=Config(
|
|
signature_version=signature_version, s3={"addressing_style": addressing_style}
|
|
),
|
|
)
|
|
return self._client
|
|
|
|
@property
|
|
def manageable(self) -> bool:
|
|
return True
|
|
|
|
def supports_file(self, name: str) -> bool:
|
|
"""We support all files"""
|
|
return True
|
|
|
|
def list_files(self) -> Generator[str]:
|
|
"""List all files returning relative paths from base_path."""
|
|
paginator = self.client.get_paginator("list_objects_v2")
|
|
pages = paginator.paginate(Bucket=self.bucket_name, Prefix=f"{self.base_path}/")
|
|
|
|
for page in pages:
|
|
for obj in page.get("Contents", []):
|
|
key = obj["Key"]
|
|
# Remove base path prefix to get relative path
|
|
rel_path = key.removeprefix(f"{self.base_path}/")
|
|
if rel_path: # Skip if it's just the directory itself
|
|
yield rel_path
|
|
|
|
def file_url(
|
|
self,
|
|
name: str,
|
|
request: HttpRequest | None = None,
|
|
use_cache: bool = True,
|
|
) -> str:
|
|
"""Generate presigned URL for file access."""
|
|
use_https = CONFIG.get_bool(
|
|
f"storage.{self.usage.value}.{self.name}.secure_urls",
|
|
CONFIG.get_bool(f"storage.{self.name}.secure_urls", True),
|
|
)
|
|
|
|
expires_in = int(
|
|
timedelta_from_string(
|
|
CONFIG.get(
|
|
f"storage.{self.usage.value}.{self.name}.url_expiry",
|
|
CONFIG.get(f"storage.{self.name}.url_expiry", "minutes=15"),
|
|
)
|
|
).total_seconds()
|
|
)
|
|
|
|
def _file_url(name: str, request: HttpRequest | None) -> str:
|
|
client = self.client
|
|
params = {
|
|
"Bucket": self.bucket_name,
|
|
"Key": f"{self.base_path}/{name}",
|
|
}
|
|
|
|
operation_name = "GetObject"
|
|
operation_model = client.meta.service_model.operation_model(operation_name)
|
|
request_dict = client._convert_to_request_dict(
|
|
params,
|
|
operation_model,
|
|
endpoint_url=client.meta.endpoint_url,
|
|
context={"is_presign_request": True},
|
|
)
|
|
|
|
# Support custom domain for S3-compatible storage (so not AWS)
|
|
# Well, can't you do custom domains on AWS as well?
|
|
custom_domain = CONFIG.get(
|
|
f"storage.{self.usage.value}.{self.name}.custom_domain",
|
|
CONFIG.get(f"storage.{self.name}.custom_domain", None),
|
|
)
|
|
if custom_domain:
|
|
scheme = "https" if use_https else "http"
|
|
path = request_dict["url_path"]
|
|
|
|
# When using path-style addressing, the presigned URL contains the bucket
|
|
# name in the path (e.g., /bucket-name/key). Since custom_domain must
|
|
# include the bucket name (per docs), strip it from the path to avoid
|
|
# duplication. See: https://github.com/goauthentik/authentik/issues/19521
|
|
# Check with trailing slash to ensure exact bucket name match
|
|
if path.startswith(f"/{self.bucket_name}/"):
|
|
path = path.removeprefix(f"/{self.bucket_name}")
|
|
|
|
# Normalize to avoid double slashes
|
|
custom_domain = custom_domain.rstrip("/")
|
|
if not path.startswith("/"):
|
|
path = f"/{path}"
|
|
|
|
custom_base = urlsplit(f"{scheme}://{custom_domain}")
|
|
|
|
# Sign the final public URL instead of signing the internal S3 endpoint and
|
|
# rewriting it afterwards. Presigned SigV4 URLs include the host header in the
|
|
# canonical request, so post-sign host changes break strict backends like RustFS.
|
|
public_path = f"{custom_base.path.rstrip('/')}{path}" if custom_base.path else path
|
|
request_dict["url_path"] = public_path
|
|
request_dict["url"] = urlunsplit(
|
|
(custom_base.scheme, custom_base.netloc, public_path, "", "")
|
|
)
|
|
|
|
return client._request_signer.generate_presigned_url(
|
|
request_dict,
|
|
operation_name,
|
|
expires_in=expires_in,
|
|
)
|
|
|
|
if use_cache:
|
|
return self._cache_get_or_set(name, request, _file_url, expires_in)
|
|
else:
|
|
return _file_url(name, request)
|
|
|
|
def save_file(self, name: str, content: bytes) -> None:
|
|
"""Save file to S3."""
|
|
self.client.put_object(
|
|
Bucket=self.bucket_name,
|
|
Key=f"{self.base_path}/{name}",
|
|
Body=content,
|
|
ACL="private",
|
|
ContentType=get_content_type(name),
|
|
)
|
|
|
|
@contextmanager
|
|
def save_file_stream(self, name: str) -> Iterator:
|
|
"""Context manager for streaming file writes to S3."""
|
|
# Keep files in memory up to 5 MB
|
|
with SpooledTemporaryFile(max_size=5 * 1024 * 1024, suffix=".S3File") as file:
|
|
yield file
|
|
file.seek(0)
|
|
self.client.upload_fileobj(
|
|
Fileobj=file,
|
|
Bucket=self.bucket_name,
|
|
Key=f"{self.base_path}/{name}",
|
|
ExtraArgs={
|
|
"ACL": "private",
|
|
"ContentType": get_content_type(name),
|
|
},
|
|
)
|
|
|
|
def delete_file(self, name: str) -> None:
|
|
"""Delete file from S3."""
|
|
self.client.delete_object(
|
|
Bucket=self.bucket_name,
|
|
Key=f"{self.base_path}/{name}",
|
|
)
|
|
|
|
def file_exists(self, name: str) -> bool:
|
|
"""Check if a file exists in S3."""
|
|
try:
|
|
self.client.head_object(
|
|
Bucket=self.bucket_name,
|
|
Key=f"{self.base_path}/{name}",
|
|
)
|
|
return True
|
|
except ClientError:
|
|
return False
|