import os
import re
import tempfile
import uuid
from dataclasses import dataclass, field
from datetime import UTC, datetime
from functools import cached_property
from pathlib import Path
from typing import TYPE_CHECKING, Any
from urllib.parse import urlsplit
import boto3
from fastapi import FastAPI, HTTPException, Request, Response, status
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel
from sqlalchemy import (
BigInteger,
DateTime,
Text,
delete,
insert,
update,
)
from sqlalchemy.orm import Mapped, mapped_column
if TYPE_CHECKING:
from collections.abc import Callable
from sqlalchemy.ext.asyncio import AsyncSession
DEFAULT_PRESIGN_TTL = 900
def _attachment_disposition(filename: str) -> str:
return f'attachment; filename="{filename}"'
def _sanitize_filename(filename: str) -> str:
base = filename.replace("\\", "/").rsplit("/", 1)[-1]
return re.sub(r"[^A-Za-z0-9._-]+", "_", base)
class StorageColumnsMixin:
s3_key: Mapped[str | None] = mapped_column(Text, nullable=True, unique=True)
content_type: Mapped[str | None] = mapped_column(Text, nullable=True)
size_bytes: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
class FileMixin(StorageColumnsMixin):
if TYPE_CHECKING:
id: Mapped[uuid.UUID]
s3_key: Mapped[str] = mapped_column(Text, unique=True)
original_filename: Mapped[str | None] = mapped_column(
Text,
nullable=True,
)
uploaded_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True),
nullable=True,
)
[docs]
@dataclass
class S3Storage:
bucket: str
region: str | None = None
endpoint_url: str | None = None
client_factory: Callable[..., Any] = field(default=boto3.client)
@cached_property
def client(self) -> Any: # noqa: ANN401
kwargs: dict[str, Any] = {"service_name": "s3"}
if self.region is not None:
kwargs["region_name"] = self.region
if self.endpoint_url is not None:
kwargs["endpoint_url"] = self.endpoint_url
return self.client_factory(**kwargs)
def presigned_put_url(
self,
key: str,
*,
expires_in: int = DEFAULT_PRESIGN_TTL,
content_type: str | None = None,
) -> str:
params: dict[str, Any] = {"Bucket": self.bucket, "Key": key}
if content_type is not None:
params["ContentType"] = content_type
url = self.client.generate_presigned_url(
"put_object",
Params=params,
ExpiresIn=expires_in,
)
return str(url)
def presigned_get_url(
self,
key: str,
*,
expires_in: int = DEFAULT_PRESIGN_TTL,
filename: str | None = None,
) -> str:
params: dict[str, Any] = {"Bucket": self.bucket, "Key": key}
if filename is not None:
params["ResponseContentDisposition"] = _attachment_disposition(
filename,
)
url = self.client.generate_presigned_url(
"get_object",
Params=params,
ExpiresIn=expires_in,
)
return str(url)
def put_bytes(
self,
key: str,
*,
blob: bytes,
content_type: str,
) -> None:
self.client.put_object(
Bucket=self.bucket,
Key=key,
Body=blob,
ContentType=content_type,
)
def delete(self, key: str) -> None:
self.client.delete_object(Bucket=self.bucket, Key=key)
[docs]
@dataclass
class LocalStorage:
root: Path
base_url: str = "/_files"
def presigned_put_url(
self,
key: str,
*,
expires_in: int = DEFAULT_PRESIGN_TTL, # noqa: ARG002 -- no TTL local
content_type: str | None = None, # noqa: ARG002 -- from PUT header
) -> str:
return f"{self.base_url}/{key}"
def presigned_get_url(
self,
key: str,
*,
expires_in: int = DEFAULT_PRESIGN_TTL, # noqa: ARG002 -- no TTL local
filename: str | None = None, # noqa: ARG002 -- static mount, no header
) -> str:
return f"{self.base_url}/{key}"
def put_bytes(
self,
key: str,
*,
blob: bytes,
content_type: str, # noqa: ARG002 -- type is served by extension
) -> None:
path = self.root / key
path.parent.mkdir(parents=True, exist_ok=True)
path.write_bytes(blob)
def delete(self, key: str) -> None:
(self.root / key).unlink(missing_ok=True)
def default_storage() -> S3Storage | LocalStorage:
storage = _resolve_storage()
if storage is None:
msg = (
"No object storage configured: set FSH_S3_BUCKET "
"(production / MinIO), or opt into local-disk dev storage "
"with FSH_LOCAL_STORAGE_DIR / FSH_LOCAL_STORAGE_URL."
)
raise RuntimeError(msg)
return storage
def _resolve_storage() -> S3Storage | LocalStorage | None:
bucket = os.environ.get("FSH_S3_BUCKET")
if bucket:
return S3Storage(
bucket=bucket,
region=os.environ.get("FSH_S3_REGION"),
endpoint_url=os.environ.get("FSH_S3_ENDPOINT_URL"),
)
dir_env = os.environ.get("FSH_LOCAL_STORAGE_DIR")
url_env = os.environ.get("FSH_LOCAL_STORAGE_URL")
if not (dir_env or url_env):
return None
root = (
Path(dir_env)
if dir_env
else Path(tempfile.gettempdir()) / "fsh_storage"
)
return LocalStorage(root=root, base_url=url_env or "/_files")
def mount_local_storage(app: FastAPI) -> None:
storage = _resolve_storage()
if not isinstance(storage, LocalStorage):
return
storage.root.mkdir(parents=True, exist_ok=True)
mount_path = urlsplit(storage.base_url).path or "/"
@app.put(f"{mount_path.rstrip('/')}/{{key:path}}", status_code=204)
async def _put_local_file(key: str, request: Request) -> Response:
storage.put_bytes(
key,
blob=await request.body(),
content_type=request.headers.get(
"content-type", "application/octet-stream"
),
)
return Response(status_code=status.HTTP_204_NO_CONTENT)
app.mount(
mount_path,
StaticFiles(directory=str(storage.root)),
name="fsh-local-files",
)
# --- Action request/response schemas --------------------------------------
[docs]
class UploadRequest(BaseModel):
filename: str
content_type: str
size_bytes: int
[docs]
class UploadResponse(BaseModel):
id: uuid.UUID
upload_url: str
[docs]
class DownloadResponse(BaseModel):
download_url: str
# --- Action functions -----------------------------------------------------
async def request_upload(
*,
model_cls: type[FileMixin],
db: AsyncSession,
auth: Any, # noqa: ANN401, ARG001 -- action handler passes this
body: UploadRequest,
) -> UploadResponse:
safe_filename = _sanitize_filename(body.filename)
key = f"{uuid.uuid4().hex}_{safe_filename}"
file_id = (
await db.execute(
insert(model_cls)
.values(
s3_key=key,
content_type=body.content_type,
size_bytes=body.size_bytes,
original_filename=safe_filename,
)
.returning(model_cls.id)
)
).scalar_one()
storage = default_storage()
upload_url = storage.presigned_put_url(
key,
content_type=body.content_type,
)
return UploadResponse(
id=file_id,
upload_url=upload_url,
)
async def complete_upload(
file: FileMixin,
*,
db: AsyncSession,
auth: Any, # noqa: ANN401, ARG001 -- action handler passes this
) -> None:
cls = type(file)
await db.execute(
update(cls)
.where(cls.id == file.id)
.values(uploaded_at=datetime.now(tz=UTC)),
)
async def download(
file: FileMixin,
*,
db: AsyncSession, # noqa: ARG001 -- action handler passes this
auth: Any, # noqa: ANN401, ARG001 -- action handler passes this
) -> DownloadResponse:
if file.uploaded_at is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="File upload not complete",
)
storage = default_storage()
return DownloadResponse(
download_url=storage.presigned_get_url(
file.s3_key,
filename=file.original_filename,
),
)
async def delete_file(
file: FileMixin,
*,
db: AsyncSession,
auth: Any, # noqa: ANN401, ARG001 -- action handler passes this
) -> None:
storage = default_storage()
storage.delete(file.s3_key)
cls = type(file)
await db.execute(delete(cls).where(cls.id == file.id))