Files
J621/backend/apps/library/services.py
T
JakeBreath 7c2569522f Make thumbnail generation atomic and warm it on import
- write thumbnails to a .part file and os.replace() them, so concurrent
  requests never read a half-written JPEG
- a stale thumbnail plus a vanished source no longer raises through the
  request (getmtime on a missing file returned 500); it falls back cleanly
- ensure_thumbnail(item) warms the preview when a file is indexed, keeping
  image decoding out of the request path
2026-09-23 21:35:09 -05:00

590 lines
20 KiB
Python

import hashlib
import json
import logging
import mimetypes
import os
import re
import shutil
import subprocess
import uuid
from datetime import datetime, timezone
from pathlib import Path
from urllib.parse import urlencode
import imagehash
from django.conf import settings
from django.http import FileResponse, Http404, HttpResponse
from django.utils.cache import get_conditional_response
from django.utils.http import http_date
from django.utils.text import get_valid_filename
from PIL import Image, ImageOps
from .models import MediaItem, MediaLocation
from .signing_urls import sign_payload
logger = logging.getLogger(__name__)
HASH_FIELDS = ("ahash", "dhash", "phash", "whash")
IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".apng", ".webp"}
ALLOWED_EXTENSIONS = {
".jpg",
".jpeg",
".png",
".gif",
".apng",
".webp",
".mp4",
".webm",
}
VIDEO_EXTENSIONS = {".mp4", ".webm"}
UPLOAD_FILE_SALT = "j621.upload-file"
MEDIA_FILE_SALT = "j621.media-file"
CHUNK_SIZE = 1024 * 1024
RANGE_RE = re.compile(r"bytes=(\d*)-(\d*)$")
# Versioned media URLs are immutable, so they may sit in the browser cache for
# as long as the signature is guaranteed to stay valid (7 days).
MEDIA_CACHE_SECONDS = 6 * 86400
# Staged uploads and similarity files can be deleted at any moment.
TEMP_CACHE_SECONDS = 3600
def compute_md5(path):
digest = hashlib.md5()
with open(path, "rb") as handle:
for chunk in iter(lambda: handle.read(CHUNK_SIZE), b""):
digest.update(chunk)
return digest.hexdigest()
def index_file(path, folder):
"""Index one file into MediaItem/MediaLocation.
Returns (item, created_item, location, created_location).
"""
path = Path(path)
folder = Path(folder)
stat = path.stat()
md5 = compute_md5(path)
item, created_item = MediaItem.objects.get_or_create(
md5=md5, defaults={"size": stat.st_size}
)
if not created_item and item.size != stat.st_size:
item.size = stat.st_size
item.save(update_fields=["size", "updated_at"])
location, created_location = MediaLocation.objects.update_or_create(
item=item,
path=str(path),
defaults={
"rel_path": str(path.relative_to(folder)),
"mtime": stat.st_mtime,
},
)
return item, created_item, location, created_location
def signed_media_url(item, user, action="raw", request=None):
"""Media URL that <img>/<video> tags can load for a signed-in user.
With a ``request`` the URL is absolute, so the SPA also works when it is
served from a different origin; without one it stays relative.
The ``v`` parameter is the item's MD5: it busts the browser cache exactly
when the file is replaced (the optimize flow rewrites files under the same
J-ID), which is what lets the URL be cached for days instead of re-minted
on every response.
"""
path = f"/api/files/J-{item.id}/{action}/"
params = {}
if user is not None and getattr(user, "is_authenticated", False):
params = {
"v": item.md5,
"sig": sign_payload(
{"item": item.id, "user": user.id, "action": action},
MEDIA_FILE_SALT,
),
}
if params:
path = f"{path}?{urlencode(params)}"
if request is None:
return path
return request.build_absolute_uri(path)
def rename_location_to_j_id(item, location):
"""Name a freshly indexed copy J-<id>.<ext> inside its own folder."""
path = Path(location.path)
if not path.exists():
return location
if path.stem == f"J-{item.id}":
return location
target = unique_destination(path.parent, f"J-{item.id}{path.suffix}")
path.rename(target)
parent = Path(location.rel_path).parent
location.path = str(target)
location.rel_path = (
str(parent / target.name) if str(parent) != "." else target.name
)
location.save(update_fields=["path", "rel_path"])
return location
def compute_visual_hashes(path):
"""Perceptual hashes for an image file (empty dict for other files)."""
path = Path(path)
if path.suffix.lower() not in IMAGE_EXTENSIONS:
return {}
try:
with Image.open(path) as image:
converted = image.convert("RGB")
return {
"ahash": str(imagehash.average_hash(converted, hash_size=8)),
"dhash": str(imagehash.dhash(converted, hash_size=8)),
"phash": str(imagehash.phash(converted, hash_size=8)),
"whash": str(imagehash.whash(converted, hash_size=8)),
}
except Exception: # noqa: BLE001 - hashing must never break indexing
logger.exception("Could not compute visual hashes for %s", path)
return {}
def ensure_visual_hashes(item):
"""Fill in missing perceptual hashes for a media item."""
if all(getattr(item, field) for field in HASH_FIELDS):
return item
location = item.locations.first()
if location is None:
return item
hashes = compute_visual_hashes(location.path)
if not hashes:
return item
for field, value in hashes.items():
setattr(item, field, value)
item.save(update_fields=[*hashes.keys(), "updated_at"])
return item
def parse_tags(raw):
"""Normalize a comma-separated string or JSON list into a list of tags."""
if raw is None:
return None
if isinstance(raw, (list, tuple)):
values = raw
else:
text = str(raw).strip()
if not text:
return None
if text.startswith("["):
try:
values = json.loads(text)
except json.JSONDecodeError:
values = text.split(",")
else:
values = text.split(",")
cleaned = []
for value in values:
tag = str(value).strip()
if tag and tag not in cleaned:
cleaned.append(tag[:100])
return cleaned
def unique_destination(folder, filename):
"""Return a non-existing path inside folder for the given filename."""
folder = Path(folder)
folder.mkdir(parents=True, exist_ok=True)
filename = get_valid_filename(Path(filename).name) or "upload"
candidate = folder / filename
stem, suffix = candidate.stem, candidate.suffix
counter = 1
while candidate.exists():
candidate = folder / f"{stem}-{counter}{suffix}"
counter += 1
return candidate
class RangeFileWrapper:
"""Iterate over a limited byte range of an open file."""
def __init__(self, file, length, chunk_size=64 * 1024):
self.file = file
self.remaining = length
self.chunk_size = chunk_size
def __iter__(self):
return self
def __next__(self):
if self.remaining <= 0:
raise StopIteration
data = self.file.read(min(self.chunk_size, self.remaining))
if not data:
raise StopIteration
self.remaining -= len(data)
return data
def close(self):
self.file.close()
def _apply_cache_headers(response, cache_control, etag, mtime):
response["Cache-Control"] = cache_control
response["ETag"] = etag
response["Last-Modified"] = http_date(mtime)
return response
def serve_file(request, path, download=False, *, max_age=TEMP_CACHE_SECONDS, immutable=False):
"""Serve a file with HTTP range support (needed for video seeking).
Responses carry validators (ETag/Last-Modified) and a private
``Cache-Control`` so browsers reuse media instead of re-downloading it on
every SPA poll. ``max_age``/``immutable`` are chosen by the caller: versioned
library media can be cached hard, staged files only briefly.
"""
path = Path(path)
if not path.is_file():
raise Http404
stat = path.stat()
size = stat.st_size
etag = f'W/"{size:x}-{stat.st_mtime_ns:x}"'
last_modified = datetime.fromtimestamp(stat.st_mtime, tz=timezone.utc)
conditional = get_conditional_response(
request, etag=etag, last_modified=last_modified
)
if conditional is not None:
return conditional
cache_control = f"private, max-age={int(max_age)}"
if immutable:
cache_control += ", immutable"
content_type = mimetypes.guess_type(str(path))[0] or "application/octet-stream"
range_header = request.headers.get("Range", "").strip()
if range_header:
match = RANGE_RE.match(range_header)
if match:
start_raw, end_raw = match.groups()
if start_raw == "" and end_raw:
length = min(int(end_raw), size)
start, end = size - length, size - 1
else:
start = int(start_raw or 0)
end = min(int(end_raw) if end_raw else size - 1, size - 1)
if start >= size or start > end:
response = HttpResponse(status=416)
response["Content-Range"] = f"bytes */{size}"
return response
length = end - start + 1
handle = open(path, "rb")
handle.seek(start)
response = FileResponse(
RangeFileWrapper(handle, length),
status=206,
content_type=content_type,
)
response["Content-Length"] = str(length)
response["Content-Range"] = f"bytes {start}-{end}/{size}"
response["Accept-Ranges"] = "bytes"
return _apply_cache_headers(
response, cache_control, etag, stat.st_mtime
)
response = FileResponse(
open(path, "rb"),
content_type=content_type,
as_attachment=download,
filename=path.name,
)
response["Accept-Ranges"] = "bytes"
return _apply_cache_headers(response, cache_control, etag, stat.st_mtime)
class DownloadCancelled(Exception):
"""Raised when a streamed download is cancelled by the user."""
class RemoteUrlError(ValueError):
"""The URL is not an allowed e621 media URL."""
def validate_remote_url(url):
"""Only http(s) URLs on the known e621 media hosts may be fetched.
Without this the download paths are an SSRF hole: any uploader could make
the server fetch internal addresses (127.0.0.1, LAN services, cloud
metadata) and read the response back through the library.
"""
from urllib.parse import urlparse
parsed = urlparse(str(url or "").strip())
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
raise RemoteUrlError("Only http(s) URLs can be fetched.")
if parsed.hostname not in settings.E621_MEDIA_HOSTS:
raise RemoteUrlError("That host is not an allowed e621 media host.")
return url
def open_remote(url, *, max_redirects=3, **kwargs):
"""GET an allowlisted URL, re-validating every redirect hop.
Returns a streaming ``requests`` response. Redirects are followed
manually so a hop cannot jump to an internal host.
"""
import requests
from urllib.parse import urljoin
current = url
for _ in range(max_redirects + 1):
validate_remote_url(current)
response = requests.get(current, allow_redirects=False, **kwargs)
if response.is_redirect or response.is_permanent_redirect:
location = response.headers.get("Location")
response.close()
if not location:
raise RemoteUrlError("The remote server redirected without a target.")
current = urljoin(current, location)
continue
return response
raise RemoteUrlError("Too many redirects from the remote server.")
def download_file(
url,
destination,
progress_callback=None,
should_cancel=None,
read_timeout=60,
):
"""Stream a remote file into destination (used by Download to Library).
The read timeout bounds how long a stalled connection can block the
worker: without it a hung socket would keep a job "downloading" forever
and the cancel flag could never be observed.
"""
headers = {"User-Agent": settings.USER_AGENT}
with open_remote(
url, headers=headers, stream=True, timeout=(10, read_timeout)
) as response:
response.raise_for_status()
total = int(response.headers.get("content-length") or 0)
downloaded = 0
with open(destination, "wb") as handle:
for chunk in response.iter_content(chunk_size=CHUNK_SIZE):
if should_cancel is not None and should_cancel():
raise DownloadCancelled("download cancelled")
if chunk:
handle.write(chunk)
downloaded += len(chunk)
if progress_callback is not None:
progress_callback(downloaded, total)
if progress_callback is not None:
progress_callback(downloaded, total or downloaded)
E621_DESCRIPTION_LIMIT = 20000
def trim_e621_post(post):
"""Keep a compact, size-bounded copy of an e621 post payload."""
if not isinstance(post, dict):
return None
file_data = post.get("file") or {}
relationships = post.get("relationships") or {}
score = post.get("score") or {}
categories = post.get("tags") or {}
return {
"id": post.get("id"),
"created_at": post.get("created_at"),
"rating": post.get("rating"),
"tags": {
str(category): [str(tag)[:200] for tag in tags][:500]
for category, tags in categories.items()
if isinstance(tags, list)
},
"score": {
"up": score.get("up"),
"down": score.get("down"),
"total": score.get("total"),
},
"fav_count": post.get("fav_count"),
"comment_count": post.get("comment_count"),
"sources": [
str(source)[:500]
for source in (post.get("sources") or [])
if source
][:20],
"description": str(post.get("description") or "")[:E621_DESCRIPTION_LIMIT],
"pools": [
int(pool)
for pool in (post.get("pools") or [])
if str(pool).isdigit()
][:50],
"relationships": {
"parent_id": relationships.get("parent_id"),
"has_children": relationships.get("has_children"),
},
"file": {
"md5": file_data.get("md5"),
"ext": file_data.get("ext"),
"size": file_data.get("size"),
"width": file_data.get("width"),
"height": file_data.get("height"),
"url": file_data.get("url"),
},
"preview": {
"url": (post.get("preview") or {}).get("url"),
},
"uploader_name": post.get("uploader_name"),
}
def sanitize_iqdb_results(results):
"""Keep a compact, size-bounded copy of IQDB candidates."""
cleaned = []
for result in results[:10]:
if not isinstance(result, dict):
continue
score = result.get("score")
post_id = result.get("post_id")
tags_preview = result.get("tags_preview")
entry = {
"post_id": post_id if isinstance(post_id, int) else None,
"score": float(score) if isinstance(score, (int, float)) else None,
"preview_url": str(result.get("preview_url") or "")[:500] or None,
"rating": str(result.get("rating") or "")[:1] or None,
"md5": str(result.get("md5") or "")[:32] or None,
"score_total": (
result.get("score_total")
if isinstance(result.get("score_total"), int)
else None
),
"fav_count": (
result.get("fav_count")
if isinstance(result.get("fav_count"), int)
else None
),
"width": (
result.get("width")
if isinstance(result.get("width"), int)
else None
),
"height": (
result.get("height")
if isinstance(result.get("height"), int)
else None
),
"tags_preview": (
[str(tag)[:100] for tag in tags_preview][:8]
if isinstance(tags_preview, list)
else []
),
}
if entry["post_id"] or entry["preview_url"]:
cleaned.append(entry)
return cleaned
def _thumbnail_is_fresh(target, path):
"""True when the cached thumbnail exists and is at least as new as source."""
try:
stat = target.stat()
if stat.st_size <= 0:
return False
return stat.st_mtime >= os.path.getmtime(path)
except OSError:
return False
def _thumbs_dir():
thumbs_dir = Path(settings.MEDIA_ROOT) / "thumbs"
thumbs_dir.mkdir(parents=True, exist_ok=True)
return thumbs_dir
def generate_video_thumbnail(md5, path):
"""Extract a JPEG thumbnail from a video, cached under MEDIA_ROOT/thumbs."""
if not shutil.which("ffmpeg"):
return None
try:
thumbs_dir = _thumbs_dir()
except OSError:
logger.exception("Could not create the thumbnail folder")
return None
target = thumbs_dir / f"{md5}.jpg"
if _thumbnail_is_fresh(target, path):
return target
# Write beside the target and move it into place, so a concurrent request
# can never read a half-written JPEG.
temp = thumbs_dir / f".{md5}.{uuid.uuid4().hex}.part.jpg"
command = [
"ffmpeg",
"-y",
"-ss",
"0.5",
"-i",
str(path),
"-frames:v",
"1",
"-vf",
"scale=480:-2",
"-loglevel",
"error",
str(temp),
]
try:
subprocess.run(command, check=True, capture_output=True, timeout=60)
os.replace(temp, target)
except (subprocess.SubprocessError, OSError):
logger.exception("Could not build a video thumbnail for %s", path)
temp.unlink(missing_ok=True)
return None
return target if target.exists() else None
def generate_image_thumbnail(md5, path):
"""Downscale an image, cached under MEDIA_ROOT/thumbs like video thumbs.
The thumbnail action used to serve full-size originals for images; a
cached 480px JPEG keeps the library grid light without touching the
original file. Returns ``None`` when the source is missing or Pillow
cannot decode it, so callers can fall back to the original.
"""
try:
thumbs_dir = _thumbs_dir()
except OSError:
logger.exception("Could not create the thumbnail folder")
return None
target = thumbs_dir / f"{md5}.jpg"
if _thumbnail_is_fresh(target, path):
return target
temp = thumbs_dir / f".{md5}.{uuid.uuid4().hex}.part.jpg"
try:
with Image.open(path) as image:
# Animated formats: the first frame is the preview.
image.seek(0)
frame = ImageOps.exif_transpose(image) or image
frame = frame.convert("RGB")
frame.thumbnail((480, 480))
frame.save(temp, "JPEG", quality=82, optimize=True)
os.replace(temp, target)
except Exception: # noqa: BLE001 - previews must never break serving
logger.exception("Could not build an image thumbnail for %s", path)
temp.unlink(missing_ok=True)
return None
return target if target.exists() else None
def ensure_thumbnail(item):
"""Generate an item's cached thumbnail if it is missing or stale.
Warming thumbnails when a file is indexed keeps image decoding out of the
request path, where the upload pipeline's hashing used to starve it.
"""
location = item.locations.first()
if location is None:
return None
path = Path(location.path)
if not path.is_file():
return None
if path.suffix.lower() in VIDEO_EXTENSIONS:
return generate_video_thumbnail(item.md5, path)
return generate_image_thumbnail(item.md5, path)