Files
J621/backend/apps/library/similarity.py
T
JakeBreath 16907c39ca Support cross-origin frontends alongside same-origin setups
- django-cors-headers with env-driven CORS_ALLOWED_ORIGINS,
  CORS_ALLOW_ALL_ORIGINS, CORS_ALLOW_CREDENTIALS and CSRF_TRUSTED_ORIGINS;
  same-origin traffic is unaffected and a disallowed origin gets no CORS
  headers. Token auth needs no cookies, so credentials stay off by default.
- TRUST_PROXY_HEADERS=true lets a TLS-terminating proxy supply
  X-Forwarded-Proto/Host for correct absolute URLs.
- API media URLs (raw/thumbnail/upload/similarity/staged previews) are now
  absolute, built from the request host, so <img>/<video>/fetch() keep
  working when the SPA is served from another origin. Signed URLs are still
  per-user; nothing is stored in the DB.
- The SPA gains VITE_API_BASE (build-time, empty = same-origin) applied by
  a small apiUrl() helper used for XHR/fetch and the few URL fallbacks.

Verified with a throwaway instance: preflight and GET responses carry the
allowed origin, foreign origins get nothing, media GETs include CORS for
cross-origin fetch(), and payload URLs use the request host (dev :8000
unchanged).
2026-09-17 22:50:12 -05:00

162 lines
5.6 KiB
Python

"""Ephemeral similarity checks.
A file lands in the media temp folder only long enough to answer "is this
already in the library, and what does it look like on e621?". Nothing is
indexed: files are deleted on startup, after the TTL, on request through the
API, and by `manage.py cleanup_similarity` (for cron).
"""
import logging
from datetime import timedelta
from pathlib import Path
from django.conf import settings
from django.core import signing
from django.utils import timezone
from rest_framework import mixins, status, viewsets
from rest_framework.decorators import action
from rest_framework.parsers import FormParser, JSONParser, MultiPartParser
from rest_framework.permissions import AllowAny, IsAuthenticated
from rest_framework.response import Response
from . import services
from .models import MediaItem, SimilarityCheck
from .serializers import SimilarityCheckSerializer
from .tools import item_brief
from .uploads import find_library_matches
logger = logging.getLogger(__name__)
TTL_MINUTES = int(getattr(settings, "SIMILARITY_TTL_MINUTES", 30))
def purge_similarity_files():
"""Delete every stored temp file without touching the database.
Safe to call during app startup (``AppConfig.ready``) — database rows are
removed lazily by ``purge_expired`` and the cleanup command.
"""
root = Path(settings.MEDIA_ROOT) / "similarity"
if not root.exists():
return 0
deleted = 0
for path in sorted(root.rglob("*")):
if not path.is_file():
continue
try:
path.unlink()
deleted += 1
except OSError as exc: # noqa: PERF203 - keep going past locked files
logger.warning("Could not delete %s: %s", path, exc)
for path in sorted(root.rglob("*"), reverse=True):
if path.is_dir():
try:
path.rmdir()
except OSError:
pass
return deleted
def purge_similarity_checks(older_than=None):
"""Delete checks (rows and files). ``older_than=None`` wipes them all."""
queryset = SimilarityCheck.objects.all()
if older_than is not None:
queryset = queryset.filter(created_at__lt=older_than)
deleted = 0
for check in queryset.iterator():
if check.file:
check.file.delete(save=False)
check.delete()
deleted += 1
return deleted
def purge_expired():
"""Remove checks past the TTL; called lazily before creating new ones."""
return purge_similarity_checks(timezone.now() - timedelta(minutes=TTL_MINUTES))
class SimilarityCheckViewSet(
mixins.ListModelMixin,
mixins.RetrieveModelMixin,
mixins.DestroyModelMixin,
viewsets.GenericViewSet,
):
"""Check an upload against the library — exact MD5 plus visual matches."""
serializer_class = SimilarityCheckSerializer
permission_classes = [IsAuthenticated]
parser_classes = [MultiPartParser, FormParser, JSONParser]
http_method_names = ["get", "post", "delete", "head", "options"]
def get_queryset(self):
queryset = SimilarityCheck.objects.all()
user = self.request.user
if not (user.is_staff or user.is_superuser):
queryset = queryset.filter(user=user)
return queryset
def create(self, request):
upload = request.FILES.get("file")
if upload is None:
return Response(
{"detail": "A file is required."}, status=status.HTTP_400_BAD_REQUEST
)
extension = Path(upload.name).suffix.lower()
if extension not in services.ALLOWED_EXTENSIONS:
return Response(
{"detail": f"Unsupported file type: {extension or 'unknown'}"},
status=status.HTTP_400_BAD_REQUEST,
)
purge_expired()
check = SimilarityCheck.objects.create(
user=request.user,
file=upload,
original_filename=upload.name,
size=upload.size,
)
check.md5 = services.compute_md5(check.file.path)
exact = MediaItem.objects.filter(md5=check.md5).first()
matches = find_library_matches(
check.file.path, limit=12, user=request.user, request=request
)
if exact is not None:
exact_j_id = f"J-{exact.id}"
matches = [
match for match in matches if match.get("j_id") != exact_j_id
]
check.results = {
"exact": item_brief(exact, request) if exact is not None else None,
"matches": matches,
}
check.save(update_fields=["md5", "results"])
return Response(
self.get_serializer(check).data, status=status.HTTP_201_CREATED
)
@action(detail=True, methods=["get", "head"], permission_classes=[AllowAny])
def file(self, request, pk=None):
"""Serve the temp file; accepts a signed URL like staged uploads."""
check = None
if request.user.is_authenticated:
check = self.get_queryset().filter(pk=pk).first()
else:
signature = request.query_params.get("sig")
payload = None
if signature:
try:
payload = signing.loads(
signature, salt=services.UPLOAD_FILE_SALT, max_age=86400
)
except signing.BadSignature:
payload = None
if payload and payload.get("check") == str(pk):
check = SimilarityCheck.objects.filter(pk=pk).first()
if check is None or not check.file:
return Response(
{"detail": "Not found."}, status=status.HTTP_404_NOT_FOUND
)
return services.serve_file(request, check.file.path)