Source code for fsh_lib.files

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))