"""Store FileField contents in Postgres (BinaryField), not on the container filesystem.""" from __future__ import annotations import mimetypes from io import BytesIO from django.core.files.base import ContentFile from django.core.files.storage import Storage from django.db import transaction from django.utils.deconstruct import deconstructible @deconstructible class DatabaseStorage(Storage): """Django storage backend backed by chat_backend.StoredFile rows.""" def _model(self): from chat_backend.models import StoredFile return StoredFile def _open(self, name, mode="rb"): stored = self._model().objects.get(name=name) return ContentFile(bytes(stored.content), name=name) def _save(self, name, content): name = self.get_available_name(name) if hasattr(content, "chunks"): data = b"".join(chunk for chunk in content.chunks()) else: data = content.read() if isinstance(data, str): data = data.encode("utf-8") content_type = getattr(content, "content_type", None) or mimetypes.guess_type(name)[0] StoredFile = self._model() with transaction.atomic(): StoredFile.objects.update_or_create( name=name, defaults={ "content": data, "size": len(data), "content_type": content_type or "", }, ) return name def delete(self, name): self._model().objects.filter(name=name).delete() def exists(self, name): return self._model().objects.filter(name=name).exists() def listdir(self, path): prefix = path.rstrip("/") if prefix: prefix = f"{prefix}/" names = self._model().objects.filter(name__startswith=prefix).values_list( "name", flat=True ) dirs: set[str] = set() files: list[str] = [] for full in names: rest = full[len(prefix) :] if prefix else full if "/" in rest: dirs.add(rest.split("/", 1)[0]) elif rest: files.append(rest) return list(dirs), files def size(self, name): return self._model().objects.values_list("size", flat=True).get(name=name) def url(self, name): # Files live in DB; serve via authenticated API / serializer when needed. return f"/api/stored-files/{name}" def path(self, name): raise NotImplementedError( "DatabaseStorage has no filesystem path; use .open()/.read() or a temp file." ) def get_accessed_time(self, name): raise NotImplementedError("DatabaseStorage does not track accessed time.") def get_created_time(self, name): return self._model().objects.values_list("created", flat=True).get(name=name) def get_modified_time(self, name): return self._model().objects.values_list("last_modified", flat=True).get(name=name)