diff --git a/ROADMAP.md b/ROADMAP.md index ec6e099..5c70504 100644 --- a/ROADMAP.md +++ b/ROADMAP.md @@ -95,18 +95,21 @@ Files now stage first and are resolved before entering the library. - [x] `cleanup_temp_uploads` command for old staged files - [x] **Auto-upload / auto-match** - [x] MD5 computed on staging; exact duplicates resolve immediately - - [x] MD5 batch-checked against e621; matches auto-complete with post - metadata stored and rating seeded -- [x] **IQDB similarity on upload** (SPA-driven) - - [x] Automatic + manual IQDB checks with candidate posts - - [x] "Visual Similarity Detected" state with candidate picker - - [x] Perceptual-hash comparison against the library (staged uploads are - flagged with their library matches as soon as they land) + - [x] MD5 batch-checked against e621 in chunks of 75; matches auto-complete + with post metadata stored and rating seeded +- [x] **Background pipeline** (server-side) + - [x] Daemon-thread worker runs MD5 → visual similarity → IQDB for every + staged upload, so the work continues after the page or tab is closed + - [x] Durable progress (`UploadRun` + per-file phase flags) polled by the + shell indicator; rate-limited runs retry with backoff + - [x] Perceptual-hash comparison against the library, loaded once per batch + - [x] IQDB candidates stored with one batched enrichment request - [x] **Upload UI** - [x] Three-column board: Pending & Unmatched / Visual Similarity Detected / - Auto-uploaded & Indexed + Auto-uploaded & Indexed, private per user (staff included) - [x] Metadata modal (link to e621 post, IQDB candidates, custom metadata) - - [x] Per-file progress plus batch processing indicator + - [x] Per-file progress plus background pipeline status + - [x] Bulk actions: bulk rate, discard all (pending/visual), dismiss all ## 4. Staff tools diff --git a/backend/apps/core/tests/test_security.py b/backend/apps/core/tests/test_security.py index 87055c9..824bd6b 100644 --- a/backend/apps/core/tests/test_security.py +++ b/backend/apps/core/tests/test_security.py @@ -345,7 +345,13 @@ class PrivacyTests(SecurityTestCase): ) self.assertEqual(self.guest.get(f"/api/uploads/{temp.id}/file/").status_code, 401) self.assertNotIn( - str(temp.id), self.client_for("sec-uploader").get("/api/uploads/").content.decode() + str(temp.id), + self.client_for("sec-uploader").get("/api/uploads/").content.decode(), + ) + # The board is per-user: staff only see their own staged uploads. + self.assertNotIn( + str(temp.id), + self.client_for("sec-staff").get("/api/uploads/").content.decode(), ) def test_similarity_checks_are_private(self): diff --git a/backend/apps/library/e621.py b/backend/apps/library/e621.py index adecd5b..72907fe 100644 --- a/backend/apps/library/e621.py +++ b/backend/apps/library/e621.py @@ -1,22 +1,39 @@ """Minimal e621 API client for server-side matching and metadata refresh. The SPA talks to e621 directly for browsing; this client exists for work the -browser cannot do reliably: long batch scans, and requests tied to a library -item rather than an open page. It uses the requesting user's stored -credentials and a global throttle (e621 asks for at most two requests per -second). +browser cannot do reliably: long batch scans, staged-upload processing and +requests tied to a library item rather than an open page. It uses the +requesting user's stored credentials and a global throttle (e621 asks for at +most two requests per second, one per second sustained). + +e621's load balancer also sheds load with 429s (sometimes with an HTML +"shedding" page instead of JSON) and the IQDB endpoint has its own, much +stricter throttle. Every call therefore retries with exponential backoff and +honours ``Retry-After``; only 401/403 are treated as fatal. """ +import logging +import random import threading import time +from pathlib import Path import requests from django.conf import settings +logger = logging.getLogger(__name__) + # Seconds between requests, per process. e621 allows 2/s hard and 1/s # sustained; each gunicorn worker throttles on its own, so leave enough # headroom that combined traffic does not trip the limit. REQUEST_INTERVAL = 1.0 +# IQDB is throttled far more aggressively than the rest of the API. +IQDB_INTERVAL = 2.0 + +MAX_ATTEMPTS = 4 +BACKOFF_BASE = 2.0 +BACKOFF_CAP = 60.0 +RETRYABLE_STATUSES = {429, 500, 502, 503, 504} class E621Error(Exception): @@ -27,6 +44,14 @@ class E621NotFound(E621Error): """The requested post does not exist (HTTP 404).""" +class E621AuthError(E621Error): + """e621 rejected the stored credentials (401/403).""" + + +class E621RateLimited(E621Error): + """e621 shed load or throttled the request after every retry.""" + + _throttle_lock = threading.Lock() _last_request_at = 0.0 @@ -35,56 +60,170 @@ def credentials_configured(user): return bool(user is not None and getattr(user, "e621_configured", False)) -def _wait_for_slot(): +def _wait_for_slot(interval=REQUEST_INTERVAL): global _last_request_at with _throttle_lock: - delay = _last_request_at + REQUEST_INTERVAL - time.monotonic() + delay = _last_request_at + interval - time.monotonic() if delay > 0: time.sleep(delay) _last_request_at = time.monotonic() +def _retry_delay(attempt, response=None): + """Backoff for a retryable failure, honouring ``Retry-After``.""" + if response is not None: + retry_after = response.headers.get("Retry-After") + if retry_after: + try: + return max(float(retry_after), 1.0) + except (TypeError, ValueError): + pass + delay = min(BACKOFF_BASE * (2**attempt), BACKOFF_CAP) + return delay + random.uniform(0, delay * 0.25) + + +def _request( + user, + method, + path, + *, + params=None, + data=None, + files=None, + timeout=30, + require_auth=True, + interval=REQUEST_INTERVAL, + attempts=MAX_ATTEMPTS, +): + """One e621 call with retries; returns the parsed JSON payload. + + ``files`` may be a callable returning the multipart mapping, which is + called once per attempt: streamed uploads consume their file handle, so a + retry needs a freshly opened file. + + Raises E621NotFound for 404s, E621AuthError for 401/403 and + E621RateLimited when e621 keeps shedding/throttling after every attempt. + """ + configured = credentials_configured(user) + if require_auth and not configured: + raise E621Error("Configure your e621 credentials in Account first.") + base = (getattr(user, "e621_base_url", "") or "https://e621.net").rstrip("/") + auth = ( + (user.e621_username, user.e621_api_key_plain) if configured else None + ) + url = f"{base}{path}" + + last_error = None + for attempt in range(attempts): + request_files = files() if callable(files) else files + _wait_for_slot(interval) + response = None + try: + response = requests.request( + method, + url, + params=params, + data=data, + files=request_files, + auth=auth, + headers={"User-Agent": settings.USER_AGENT}, + timeout=timeout, + ) + except requests.RequestException as exc: + last_error = E621Error(f"Could not reach e621: {exc}") + else: + if response.status_code == 404: + raise E621NotFound(f"e621 returned 404 for {path}") + if response.status_code in {401, 403}: + raise E621AuthError( + f"e621 rejected the request ({response.status_code}). " + "Check the stored e621 credentials." + ) + if response.status_code == 429: + # Throttles and load-shedding can arrive as JSON ({"message": + # "Throttled: ..."}) or as an HTML page. + message = "" + try: + payload = response.json() + except ValueError: + payload = None + if isinstance(payload, dict): + message = str( + payload.get("message") or payload.get("error") or "" + ) + last_error = E621RateLimited( + message or f"e621 throttled the request for {path}" + ) + elif response.status_code < 400: + try: + return response.json() + except ValueError: + # An HTML page with a 2xx status. + last_error = E621RateLimited( + f"e621 returned an unexpected {response.status_code} response." + ) + elif response.status_code in RETRYABLE_STATUSES: + last_error = E621RateLimited( + f"e621 replied {response.status_code} for {path}" + ) + else: + raise E621Error(f"e621 replied {response.status_code} for {path}") + finally: + _close_upload_files(request_files) + if attempt + 1 < attempts: + delay = _retry_delay(attempt, response) + logger.info( + "e621 %s %s failed (%s); retrying in %.1fs", + method, + path, + last_error, + delay, + ) + time.sleep(delay) + + if last_error is None: + last_error = E621Error("e621 request failed.") + raise last_error + + +def _close_upload_files(files): + """Close the handles behind a multipart mapping (see _request).""" + if not isinstance(files, dict): + return + for value in files.values(): + handle = value[1] if isinstance(value, tuple) and len(value) > 1 else value + close = getattr(handle, "close", None) + if close is not None: + try: + close() + except Exception: # noqa: BLE001 - closing must never mask errors + pass + + def get(user, path, params=None, timeout=30, require_auth=True): """GET an e621 API path using the user's credentials. Reads that e621 serves anonymously (searches, pools, tags) can pass require_auth=False; matching endpoints keep requiring credentials. - - Raises E621NotFound for 404s and E621Error for everything else that isn't - a 2xx, so callers never see requests exceptions. """ - configured = credentials_configured(user) - if require_auth and not configured: - raise E621Error("Configure your e621 credentials in Account first.") - base = (getattr(user, "e621_base_url", "") or "https://e621.net").rstrip("/") - _wait_for_slot() - try: - response = requests.get( - f"{base}{path}", - params=params, - auth=( - (user.e621_username, user.e621_api_key_plain) - if configured - else None - ), - headers={"User-Agent": settings.USER_AGENT}, - timeout=timeout, - ) - except requests.RequestException as exc: - raise E621Error(f"Could not reach e621: {exc}") from exc - if response.status_code == 404: - raise E621NotFound(f"e621 returned 404 for {path}") - if response.status_code >= 400: - raise E621Error(f"e621 replied {response.status_code} for {path}") - try: - return response.json() - except ValueError as exc: - raise E621Error("e621 returned an unexpected response.") from exc + return _request( + user, + "GET", + path, + params=params, + timeout=timeout, + require_auth=require_auth, + ) def find_post_by_md5(user, md5): """The e621 post with this exact MD5, or None.""" - payload = get(user, "/posts.json", params={"tags": f"md5:{md5}", "limit": 1}) + payload = get( + user, + "/posts.json", + params={"tags": f"md5:{md5}", "limit": 1}, + require_auth=False, + ) posts = payload.get("posts") if isinstance(payload, dict) else None if not posts: return None @@ -98,3 +237,87 @@ def fetch_post(user, post_id): if not isinstance(post, dict): raise E621Error("e621 returned an unexpected post payload.") return post + + +def check_md5_batch(user, md5s): + """Look many MD5s up in one posts.json query. + + Returns ``{md5: post}`` for the ones e621 knows; missing MD5s are simply + absent. Works anonymously, like the original app's batch cache command. + """ + wanted = {str(value).strip().lower() for value in md5s if value} + if not wanted: + return {} + values = sorted(wanted) + payload = get( + user, + "/posts.json", + params={ + "tags": f"md5:{','.join(values)}", + "limit": min(len(values), 320), + }, + require_auth=False, + ) + posts = payload.get("posts") if isinstance(payload, dict) else None + found = {} + for post in posts or []: + if not isinstance(post, dict): + continue + file_data = post.get("file") or {} + md5 = str(file_data.get("md5") or "").strip().lower() + if md5 in wanted: + found[md5] = post + return found + + +def fetch_posts_by_ids(user, ids): + """Fetch many posts in one query (up to 320 ids). Missing ids are absent.""" + values = sorted({int(value) for value in ids}) + if not values: + return [] + payload = get( + user, + "/posts.json", + params={ + "tags": f"id:{','.join(str(value) for value in values)}", + "limit": min(len(values), 320), + }, + require_auth=False, + ) + posts = payload.get("posts") if isinstance(payload, dict) else None + return [post for post in posts or [] if isinstance(post, dict)] + + +def iqdb_search(user, path, timeout=60): + """Reverse-image search one file through e621's IQDB endpoint. + + Returns the legacy match list. Uses the extra-strict IQDB interval and + retries through e621's throttle; raises E621RateLimited when it persists. + """ + path = Path(path) + + def open_file(): + # A fresh handle per attempt: the stream is consumed by the request. + return {"search[file]": (path.name, open(path, "rb"))} + + payload = _request( + user, + "POST", + "/iqdb_queries.json", + files=open_file, + timeout=timeout, + require_auth=False, + interval=IQDB_INTERVAL, + ) + if isinstance(payload, list): + return payload + if isinstance(payload, dict): + matches = payload.get("matches") + if isinstance(matches, list): + return matches + # e621 answers its throttle with {"success": false, "message": ...}. + message = payload.get("message") or payload.get("error") + if message: + raise E621RateLimited(str(message)) + raise E621Error("e621 returned an unexpected IQDB payload.") + return [] diff --git a/backend/apps/library/management/commands/cleanup_temp_uploads.py b/backend/apps/library/management/commands/cleanup_temp_uploads.py index 47c309d..9340ec9 100644 --- a/backend/apps/library/management/commands/cleanup_temp_uploads.py +++ b/backend/apps/library/management/commands/cleanup_temp_uploads.py @@ -18,6 +18,9 @@ class Command(BaseCommand): ) def handle(self, *args, **options): + from apps.library.upload_pipeline import reap_stale_claims + + reap_stale_claims() cutoff = timezone.now() - timedelta(hours=options["hours"]) queryset = TempUpload.objects.filter(created_at__lt=cutoff) if not options["include_completed"]: diff --git a/backend/apps/library/management/commands/process_uploads.py b/backend/apps/library/management/commands/process_uploads.py new file mode 100644 index 0000000..4a19228 --- /dev/null +++ b/backend/apps/library/management/commands/process_uploads.py @@ -0,0 +1,68 @@ +"""Process staged uploads: e621 MD5, visual similarity, IQDB. + +Runs synchronously, unlike the daemon thread the API starts on demand. Useful +for tests, for a manual drain after an outage and for the scheduler if a +deployment wants a periodic safety net. + + python manage.py process_uploads # every user with queued work, once + python manage.py process_uploads --user 3 # one user + python manage.py process_uploads --loop 60 # keep draining every 60s +""" + +import time + +from django.core.management.base import BaseCommand + +from apps.library.models import TempUpload +from apps.library.upload_pipeline import ( + MAX_ATTEMPTS, + OUTSTANDING_Q, + WORK_STATUSES, + reap_stale_claims, + run_pipeline, +) + + +class Command(BaseCommand): + help = "Run the staged-upload pipeline (MD5 -> visual similarity -> IQDB)." + + def add_arguments(self, parser): + parser.add_argument( + "--user", + type=int, + default=None, + help="Only process this user id.", + ) + parser.add_argument( + "--loop", + type=int, + default=0, + metavar="SECONDS", + help="Keep draining every SECONDS seconds instead of exiting.", + ) + + def handle(self, *args, **options): + interval = options["loop"] or 0 + while True: + self.drain(user_id=options["user"]) + if interval <= 0: + return + time.sleep(interval) + + def drain(self, user_id=None): + reap_stale_claims() + queryset = ( + TempUpload.objects.filter(status__in=WORK_STATUSES) + .filter(OUTSTANDING_Q) + .filter(attempts__lt=MAX_ATTEMPTS) + ) + if user_id is not None: + queryset = queryset.filter(user_id=user_id) + user_ids = list(queryset.values_list("user_id", flat=True).distinct()) + if not user_ids: + self.stdout.write("No staged uploads need processing.") + return + for value in user_ids: + self.stdout.write(f"Processing staged uploads for user {value}...") + run_pipeline(value) + self.stdout.write(self.style.SUCCESS(f"Processed {len(user_ids)} queue(s).")) diff --git a/backend/apps/library/migrations/0011_upload_pipeline.py b/backend/apps/library/migrations/0011_upload_pipeline.py new file mode 100644 index 0000000..7c32036 --- /dev/null +++ b/backend/apps/library/migrations/0011_upload_pipeline.py @@ -0,0 +1,81 @@ +# Generated by Django 6.1.1 on 2026-09-21 13:27 + +import django.db.models.deletion +from django.conf import settings +from django.db import migrations, models + + +def backfill_pipeline_checks(apps, schema_editor): + """Mark pre-pipeline rows as already MD5/visual-checked when they were. + + Rows with an e621 post id (or already completed) clearly went through the + MD5 phase; rows with visual matches went through the visual phase. + Everything else stays unset so the new pipeline picks it up once after + deploy — a re-check of stale pending uploads is the desired behavior. + """ + TempUpload = apps.get_model("library", "TempUpload") + TempUpload.objects.filter( + models.Q(e621_post_id__isnull=False) | models.Q(status="completed") + ).update(e621_checked_at=models.F("updated_at")) + TempUpload.objects.filter(visual_matches__isnull=False).update( + visual_checked_at=models.F("updated_at") + ) + + +class Migration(migrations.Migration): + + dependencies = [ + ('library', '0010_similaritycheck'), + migrations.swappable_dependency(settings.AUTH_USER_MODEL), + ] + + operations = [ + migrations.AddField( + model_name='tempupload', + name='attempts', + field=models.PositiveSmallIntegerField(default=0), + ), + migrations.AddField( + model_name='tempupload', + name='claimed_at', + field=models.DateTimeField(blank=True, db_index=True, null=True), + ), + migrations.AddField( + model_name='tempupload', + name='e621_checked_at', + field=models.DateTimeField(blank=True, null=True), + ), + migrations.AddField( + model_name='tempupload', + name='pipeline_error', + field=models.TextField(blank=True, default=''), + ), + migrations.AddField( + model_name='tempupload', + name='visual_checked_at', + field=models.DateTimeField(blank=True, null=True), + ), + migrations.CreateModel( + name='UploadRun', + fields=[ + ('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')), + ('status', models.CharField(choices=[('idle', 'Idle'), ('running', 'Running'), ('paused', 'Paused'), ('error', 'Error')], default='idle', max_length=20)), + ('phase', models.CharField(blank=True, choices=[('', 'None'), ('md5', 'e621 MD5'), ('visual', 'Visual similarity'), ('iqdb', 'IQDB')], default='', max_length=20)), + ('total', models.IntegerField(default=0)), + ('processed', models.IntegerField(default=0)), + ('matched', models.IntegerField(default=0)), + ('failed', models.IntegerField(default=0)), + ('error', models.TextField(blank=True, default='')), + ('started_at', models.DateTimeField(blank=True, null=True)), + ('updated_at', models.DateTimeField(auto_now=True)), + ('user', models.OneToOneField(on_delete=django.db.models.deletion.CASCADE, related_name='upload_run', to=settings.AUTH_USER_MODEL)), + ], + options={ + 'ordering': ['-updated_at'], + }, + ), + migrations.RunPython( + code=backfill_pipeline_checks, reverse_code=migrations.RunPython.noop + ), + ] + diff --git a/backend/apps/library/models.py b/backend/apps/library/models.py index 02f2c4a..c61bfc9 100644 --- a/backend/apps/library/models.py +++ b/backend/apps/library/models.py @@ -164,6 +164,16 @@ class TempUpload(models.Model): on_delete=models.SET_NULL, related_name="temp_uploads", ) + # Background pipeline bookkeeping (see apps/library/upload_pipeline.py). + # e621_checked_at/visual_checked_at are set once the corresponding phase + # ran, so "never checked" and "checked, nothing found" stay distinct. + e621_checked_at = models.DateTimeField(null=True, blank=True) + visual_checked_at = models.DateTimeField(null=True, blank=True) + # Worker claim for cross-process mutual exclusion; stale claims are + # reaped and the row queued again. + claimed_at = models.DateTimeField(null=True, blank=True, db_index=True) + pipeline_error = models.TextField(blank=True, default="") + attempts = models.PositiveSmallIntegerField(default=0) created_at = models.DateTimeField(auto_now_add=True) updated_at = models.DateTimeField(auto_now=True) @@ -174,6 +184,61 @@ class TempUpload(models.Model): return f"{self.original_filename} ({self.status})" +class UploadRun(models.Model): + """Per-user state of the background upload pipeline. + + One row per user acts as the cheap status source the SPA polls and as a + place for batch-level failures (broken e621 credentials, outages) that + would otherwise be repeated on every row. + """ + + STATUS_IDLE = "idle" + STATUS_RUNNING = "running" + STATUS_PAUSED = "paused" + STATUS_ERROR = "error" + STATUS_CHOICES = [ + (STATUS_IDLE, "Idle"), + (STATUS_RUNNING, "Running"), + (STATUS_PAUSED, "Paused"), + (STATUS_ERROR, "Error"), + ] + + PHASE_MD5 = "md5" + PHASE_VISUAL = "visual" + PHASE_IQDB = "iqdb" + PHASE_CHOICES = [ + ("", "None"), + (PHASE_MD5, "e621 MD5"), + (PHASE_VISUAL, "Visual similarity"), + (PHASE_IQDB, "IQDB"), + ] + + user = models.OneToOneField( + settings.AUTH_USER_MODEL, + on_delete=models.CASCADE, + related_name="upload_run", + ) + status = models.CharField( + max_length=20, choices=STATUS_CHOICES, default=STATUS_IDLE + ) + phase = models.CharField( + max_length=20, choices=PHASE_CHOICES, blank=True, default="" + ) + total = models.IntegerField(default=0) + processed = models.IntegerField(default=0) + matched = models.IntegerField(default=0) + failed = models.IntegerField(default=0) + error = models.TextField(blank=True, default="") + started_at = models.DateTimeField(null=True, blank=True) + updated_at = models.DateTimeField(auto_now=True) + + class Meta: + ordering = ["-updated_at"] + + def __str__(self): + return f"Upload run for {self.user_id} ({self.status})" + + class DownloadTask(models.Model): """A background 'Download to Library' job with progress tracking.""" diff --git a/backend/apps/library/serializers.py b/backend/apps/library/serializers.py index 6f5ee6c..55563db 100644 --- a/backend/apps/library/serializers.py +++ b/backend/apps/library/serializers.py @@ -145,6 +145,11 @@ class TempUploadSerializer(serializers.ModelSerializer): library_j_id = serializers.SerializerMethodField() file_url = serializers.SerializerMethodField() preview_url = serializers.SerializerMethodField() + md5_checked = serializers.SerializerMethodField() + visual_checked = serializers.SerializerMethodField() + iqdb_checked = serializers.SerializerMethodField() + processing = serializers.SerializerMethodField() + similar_count = serializers.SerializerMethodField() class Meta: model = TempUpload @@ -165,6 +170,13 @@ class TempUploadSerializer(serializers.ModelSerializer): "library_j_id", "file_url", "preview_url", + "pipeline_error", + "attempts", + "md5_checked", + "visual_checked", + "iqdb_checked", + "processing", + "similar_count", "created_at", "updated_at", ] @@ -215,6 +227,38 @@ class TempUploadSerializer(serializers.ModelSerializer): item, user, action, request=self.context.get("request") ) + def get_md5_checked(self, obj): + return obj.e621_checked_at is not None + + def get_visual_checked(self, obj): + return obj.visual_checked_at is not None + + def get_iqdb_checked(self, obj): + return obj.iqdb_data is not None + + def get_processing(self, obj): + return obj.claimed_at is not None + + def get_similar_count(self, obj): + return len(obj.iqdb_data or []) + len(obj.visual_matches or []) + + +class TempUploadListSerializer(TempUploadSerializer): + """Compact staged-upload row for the board and the status polling. + + Drops the heavy post/IQDB payloads (the metadata modal fetches the full + row) while keeping the pipeline flags the board renders per tile. + """ + + class Meta(TempUploadSerializer.Meta): + fields = [ + field + for field in TempUploadSerializer.Meta.fields + if field + not in {"e621_data", "iqdb_data", "visual_matches", "custom_tags", "custom_notes"} + ] + read_only_fields = fields + class DownloadTaskSerializer(serializers.ModelSerializer): task_id = serializers.UUIDField(source="id", read_only=True) diff --git a/backend/apps/library/tests/test_e621.py b/backend/apps/library/tests/test_e621.py new file mode 100644 index 0000000..b7d6cc6 --- /dev/null +++ b/backend/apps/library/tests/test_e621.py @@ -0,0 +1,99 @@ +"""The server-side e621 client: batching, retries and IQDB stream handling.""" + +from pathlib import Path +from tempfile import TemporaryDirectory +from unittest import mock + +from django.test import SimpleTestCase + +from apps.library import e621 + + +class FakeResponse: + def __init__(self, status_code, payload=None, headers=None): + self.status_code = status_code + self._payload = payload + self.headers = headers or {} + + def json(self): + if self._payload is None: + raise ValueError("not json") + return self._payload + + +class E621ClientTests(SimpleTestCase): + def test_iqdb_search_reopens_the_file_on_retry(self): + bodies = [] + + def fake_request(method, url, **kwargs): + handle = kwargs["files"]["search[file]"][1] + bodies.append(handle.read()) + if len(bodies) == 1: + return FakeResponse( + 429, {"success": False, "message": "Throttled"} + ) + return FakeResponse(200, [{"post_id": 1, "score": 90.0}]) + + with TemporaryDirectory() as tmp: + path = Path(tmp) / "x.png" + path.write_bytes(b"image-bytes") + with mock.patch.object( + e621.requests, "request", side_effect=fake_request + ), mock.patch.object(e621.time, "sleep"), mock.patch.object( + e621, "_wait_for_slot" + ): + result = e621.iqdb_search(None, path) + + # Both attempts must send the full body, not the consumed handle. + self.assertEqual(bodies, [b"image-bytes", b"image-bytes"]) + self.assertEqual(result, [{"post_id": 1, "score": 90.0}]) + + def test_check_md5_batch_keys_by_md5(self): + payload = { + "posts": [ + {"id": 5, "file": {"md5": "a" * 32}}, + {"id": 6, "file": {"md5": "b" * 32}}, + ] + } + with mock.patch.object(e621, "get", return_value=payload) as getter: + found = e621.check_md5_batch(None, ["A" * 32, "b" * 32]) + self.assertEqual(set(found), {"a" * 32, "b" * 32}) + params = getter.call_args.kwargs["params"] + self.assertTrue(params["tags"].startswith("md5:")) + self.assertEqual(params["limit"], 2) + + def test_auth_errors_are_not_retried(self): + calls = [] + + def fake_request(*args, **kwargs): + calls.append(1) + return FakeResponse(403, {"error": "nope"}) + + with mock.patch.object( + e621.requests, "request", side_effect=fake_request + ), mock.patch.object(e621, "_wait_for_slot"): + with self.assertRaises(e621.E621AuthError): + e621._request(None, "GET", "/posts.json", require_auth=False) + self.assertEqual(len(calls), 1) + + def test_load_shedding_html_raises_rate_limited_after_retries(self): + with mock.patch.object( + e621.requests, + "request", + return_value=FakeResponse(200, None), + ), mock.patch.object(e621.time, "sleep"), mock.patch.object( + e621, "_wait_for_slot" + ): + with self.assertRaises(e621.E621RateLimited): + e621._request( + None, "GET", "/posts.json", require_auth=False, attempts=2 + ) + + def test_fetch_posts_by_ids_queries_with_id_tag(self): + with mock.patch.object( + e621, "get", return_value={"posts": [{"id": 9}]} + ) as getter: + posts = e621.fetch_posts_by_ids(None, [9]) + self.assertEqual(posts, [{"id": 9}]) + params = getter.call_args.kwargs["params"] + self.assertEqual(params["tags"], "id:9") diff --git a/backend/apps/library/tests/test_uploads.py b/backend/apps/library/tests/test_uploads.py index ea44c7a..0dc5d11 100644 --- a/backend/apps/library/tests/test_uploads.py +++ b/backend/apps/library/tests/test_uploads.py @@ -6,17 +6,21 @@ import io import json import shutil import tempfile +from datetime import timedelta from pathlib import Path +from unittest import mock from django.contrib.auth import get_user_model from django.core.files.uploadedfile import SimpleUploadedFile from django.test import Client, TestCase, override_settings +from django.utils import timezone from PIL import Image from rest_framework.authtoken.models import Token -from apps.library.models import MediaItem, TempUpload +from apps.library import upload_pipeline +from apps.library.models import MediaItem, TempUpload, UploadRun User = get_user_model() @@ -95,6 +99,7 @@ class TempUploadListTests(TestCase): self.assertEqual(Client().get("/api/uploads/").status_code, 401) +@override_settings(UPLOAD_PIPELINE_AUTOSTART=False) class StagedUploadWorkflowTests(TestCase): @classmethod def setUpClass(cls): @@ -400,3 +405,293 @@ class IqdbRecordingTests(TestCase): self.client, f"/api/uploads/{temp.id}/iqdb/", {"results": "nope"} ) self.assertEqual(response.status_code, 400) + + +@override_settings(UPLOAD_PIPELINE_AUTOSTART=False) +class UploadPipelineTests(TestCase): + """The server-side MD5 -> visual -> IQDB queue and its board API.""" + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls._tmp = tempfile.mkdtemp(prefix="j621-pipeline-") + cls._watched = Path(cls._tmp) / "library" + cls._watched.mkdir(parents=True, exist_ok=True) + cls._settings = override_settings( + MEDIA_ROOT=cls._tmp, WATCHED_FOLDER=str(cls._watched) + ) + cls._settings.enable() + + @classmethod + def tearDownClass(cls): + cls._settings.disable() + shutil.rmtree(cls._tmp, ignore_errors=True) + super().tearDownClass() + + def setUp(self): + self.uploader = User.objects.create_user( + username="pipe-uploader", password="pipe-pass-123456" + ) + self.uploader.role = "uploader" + self.uploader.save(update_fields=["role"]) + self.other = User.objects.create_user( + username="pipe-other", password="pipe-pass-123456" + ) + self.other.role = "uploader" + self.other.save(update_fields=["role"]) + + def api_client(self, user): + client = Client() + client.defaults["HTTP_AUTHORIZATION"] = ( + f"Token {Token.objects.create(user=user).key}" + ) + return client + + def stage(self, user=None, label="file", payload=None, filename=None): + user = user or self.uploader + payload = payload or TINY_PNG + name = filename or f"{label}.png" + return TempUpload.objects.create( + user=user, + file=SimpleUploadedFile(name, payload, content_type="image/png"), + original_filename=name, + md5=hashlib.md5(payload + label.encode()).hexdigest(), + size=len(payload), + ) + + def test_md5_match_auto_imports_the_file(self): + temp = self.stage(label="match") + post = { + "id": 123456, + "rating": "s", + "file": { + "md5": temp.md5, + "url": "https://static1.e621.net/data/m.png", + }, + "tags": {"general": ["canine"]}, + } + with mock.patch.object( + upload_pipeline.e621, + "check_md5_batch", + return_value={temp.md5: post}, + ): + upload_pipeline.run_pipeline(self.uploader.id) + + temp.refresh_from_db() + self.assertEqual(temp.status, TempUpload.STATUS_COMPLETED) + self.assertEqual(temp.resolution, TempUpload.RESOLUTION_AUTO_MD5) + self.assertEqual(temp.e621_post_id, 123456) + self.assertIsNotNone(temp.library_item_id) + self.assertIsNotNone(temp.e621_checked_at) + self.assertIsNone(temp.claimed_at) + run = UploadRun.objects.get(user=self.uploader) + self.assertEqual(run.status, UploadRun.STATUS_IDLE) + self.assertEqual(run.matched, 1) + self.assertEqual(run.processed, 1) + + def test_unmatched_file_runs_every_phase(self): + temp = self.stage(label="nomatch") + raw_iqdb = [ + { + "post_id": 777, + "score": 91.0, + "post": { + "id": 777, + "rating": "q", + "md5": "a" * 32, + "score": 5, + "fav_count": 2, + "image_width": 800, + "image_height": 600, + }, + } + ] + modern = [ + { + "id": 777, + "rating": "q", + "fav_count": 4, + "score": {"total": 9}, + "preview": {"url": "https://static1.e621.net/data/preview/x.jpg"}, + "file": {"md5": "a" * 32, "width": 801, "height": 601}, + "tags": {"general": ["canine", "solo"]}, + } + ] + with mock.patch.object( + upload_pipeline.e621, "check_md5_batch", return_value={} + ), mock.patch.object( + upload_pipeline.e621, "iqdb_search", return_value=raw_iqdb + ), mock.patch.object( + upload_pipeline.e621, "fetch_posts_by_ids", return_value=modern + ): + upload_pipeline.run_pipeline(self.uploader.id) + + temp.refresh_from_db() + self.assertIsNotNone(temp.e621_checked_at) + self.assertIsNotNone(temp.visual_checked_at) + self.assertEqual(temp.status, TempUpload.STATUS_VISUAL_MATCH) + self.assertEqual(len(temp.iqdb_data), 1) + candidate = temp.iqdb_data[0] + self.assertEqual(candidate["post_id"], 777) + self.assertEqual(candidate["score_total"], 9) + self.assertEqual(candidate["fav_count"], 4) + self.assertEqual(candidate["width"], 801) + self.assertEqual( + candidate["preview_url"], "https://static1.e621.net/data/preview/x.jpg" + ) + self.assertEqual(candidate["tags_preview"], ["canine", "solo"]) + run = UploadRun.objects.get(user=self.uploader) + self.assertEqual(run.status, UploadRun.STATUS_IDLE) + self.assertEqual(run.processed, 1) + + def test_visual_match_flags_similar_library_items(self): + from apps.library.uploads import complete_temp_upload + + seed = self.stage(label="seed") + complete_temp_upload(seed) + self.assertTrue(MediaItem.objects.exists()) + + temp = self.stage(label="similar") + with mock.patch.object( + upload_pipeline.e621, "check_md5_batch", return_value={} + ), mock.patch.object( + upload_pipeline.e621, "iqdb_search", return_value=[] + ): + upload_pipeline.run_pipeline(self.uploader.id) + + temp.refresh_from_db() + self.assertGreaterEqual(len(temp.visual_matches or []), 1) + self.assertEqual(temp.status, TempUpload.STATUS_VISUAL_MATCH) + + def test_iqdb_rate_limit_pauses_the_run(self): + temp = self.stage(label="throttled") + with mock.patch.object( + upload_pipeline.e621, "check_md5_batch", return_value={} + ), mock.patch.object( + upload_pipeline.e621, + "iqdb_search", + side_effect=upload_pipeline.e621.E621RateLimited("Throttled"), + ): + upload_pipeline.run_pipeline(self.uploader.id) + + run = UploadRun.objects.get(user=self.uploader) + self.assertEqual(run.status, UploadRun.STATUS_PAUSED) + self.assertIn("e621", run.error) + temp.refresh_from_db() + self.assertEqual(temp.attempts, 1) + self.assertNotEqual(temp.pipeline_error, "") + self.assertIsNone(temp.iqdb_data) + self.assertIsNone(temp.claimed_at) + + client = self.api_client(self.uploader) + response = jpost(client, f"/api/uploads/{temp.id}/retry/", {}) + self.assertEqual(response.status_code, 200) + temp.refresh_from_db() + self.assertEqual(temp.attempts, 0) + self.assertEqual(temp.pipeline_error, "") + + def test_failed_rows_stop_after_the_attempt_cap(self): + temp = self.stage(label="broken") + for _ in range(upload_pipeline.MAX_ATTEMPTS): + with mock.patch.object( + upload_pipeline.e621, "check_md5_batch", return_value={} + ), mock.patch.object( + upload_pipeline.e621, + "iqdb_search", + side_effect=upload_pipeline.e621.E621Error("boom"), + ): + upload_pipeline.run_pipeline(self.uploader.id) + temp.refresh_from_db() + self.assertEqual(temp.attempts, upload_pipeline.MAX_ATTEMPTS) + self.assertEqual(upload_pipeline.count_outstanding(self.uploader), 0) + run = UploadRun.objects.get(user=self.uploader) + self.assertGreaterEqual(run.failed, 1) + + def test_status_process_and_compact_board_payload(self): + temp = self.stage(label="board") + client = self.api_client(self.uploader) + + status = client.get("/api/uploads/status/").json() + self.assertEqual(status["status"], UploadRun.STATUS_IDLE) + self.assertTrue(status["active"]) + self.assertEqual(status["outstanding"], 1) + self.assertEqual(status["waiting"]["md5"], 1) + + self.assertEqual(client.post("/api/uploads/process/").status_code, 200) + + rows = client.get("/api/uploads/").json() + self.assertEqual(len(rows), 1) + row = rows[0] + for key in ( + "md5_checked", + "visual_checked", + "iqdb_checked", + "processing", + "similar_count", + "pipeline_error", + ): + self.assertIn(key, row) + self.assertNotIn("e621_data", row) + self.assertNotIn("iqdb_data", row) + + detail = client.get(f"/api/uploads/{temp.id}/").json() + self.assertIn("e621_data", detail) + self.assertIn("iqdb_data", detail) + + def test_discard_bulk_removes_only_own_rows(self): + client = self.api_client(self.uploader) + first = self.stage(label="discard-a") + second = self.stage(label="discard-b") + theirs = self.stage(self.other, label="discard-theirs") + paths = [Path(first.file.path), Path(second.file.path)] + + response = jpost( + client, + "/api/uploads/discard-bulk/", + {"temp_ids": [str(first.id), str(second.id), str(theirs.id)]}, + ) + self.assertEqual(response.status_code, 200) + body = response.json() + self.assertEqual(len(body["discarded"]), 2) + self.assertEqual(len(body["errors"]), 1) + self.assertEqual(body["errors"][0]["error"], "not found") + for path in paths: + self.assertFalse(path.exists()) + self.assertFalse( + TempUpload.objects.filter(pk__in=[first.id, second.id]).exists() + ) + self.assertTrue(TempUpload.objects.filter(pk=theirs.id).exists()) + + def test_retry_with_phase_rechecks_iqdb(self): + temp = self.stage(label="recheck") + TempUpload.objects.filter(pk=temp.pk).update( + iqdb_data=[{"post_id": 1}], + status=TempUpload.STATUS_VISUAL_MATCH, + ) + client = self.api_client(self.uploader) + response = jpost( + client, f"/api/uploads/{temp.id}/retry/", {"phase": "iqdb"} + ) + self.assertEqual(response.status_code, 200) + temp.refresh_from_db() + self.assertIsNone(temp.iqdb_data) + self.assertEqual(temp.status, TempUpload.STATUS_VISUAL_MATCH) + + def test_stale_claims_are_released(self): + temp = self.stage(label="stale") + TempUpload.objects.filter(pk=temp.pk).update( + claimed_at=timezone.now() - upload_pipeline.STALE_CLAIM_AFTER + - timedelta(minutes=1) + ) + UploadRun.objects.create(user=self.uploader, status=UploadRun.STATUS_RUNNING) + UploadRun.objects.filter(user=self.uploader).update( + updated_at=timezone.now() - upload_pipeline.STALE_RUN_AFTER + - timedelta(minutes=1) + ) + released, paused = upload_pipeline.reap_stale_claims() + self.assertEqual(released, 1) + self.assertEqual(paused, 1) + temp.refresh_from_db() + self.assertIsNone(temp.claimed_at) + run = UploadRun.objects.get(user=self.uploader) + self.assertEqual(run.status, UploadRun.STATUS_PAUSED) diff --git a/backend/apps/library/upload_pipeline.py b/backend/apps/library/upload_pipeline.py new file mode 100644 index 0000000..dab02aa --- /dev/null +++ b/backend/apps/library/upload_pipeline.py @@ -0,0 +1,589 @@ +"""Background processing for staged uploads. + +Each staged upload runs through three phases, in batches of 75 (the same +lookup size the original app used for its e621 MD5 cache command): + +1. e621 MD5 lookup — one ``posts.json`` query per round; byte-identical + matches are imported straight into the library (``auto_md5``). +2. Local visual similarity — perceptual hashes are compared against the + whole library once per round. +3. e621 IQDB — reverse-image search for whatever is still unresolved. + +The pipeline runs in a daemon thread started on demand (like the download and +match scans), so the browser can navigate away and the work keeps going. +Progress and pause/error state live in the ``UploadRun`` row; per-file state +lives on ``TempUpload``. Rows are claimed with ``SELECT ... FOR UPDATE SKIP +LOCKED`` so several gunicorn workers cannot process the same file, and stale +claims left by a recycled worker are reaped and picked up again. +""" + +import logging +import threading +import time +from datetime import timedelta + +from django.conf import settings +from django.contrib.auth import get_user_model +from django.db import connection, transaction +from django.db.models import F, Q +from django.utils import timezone + +from . import e621, services +from .models import TempUpload, UploadRun + +logger = logging.getLogger(__name__) + +# Same batch size as the original app's e621 cache command. +MD5_BATCH_SIZE = 75 +# Rows one worker round claims; also the MD5 query size. +CLAIM_SIZE = MD5_BATCH_SIZE +# A claimed row is assumed dead after this long and is queued again. +STALE_CLAIM_AFTER = timedelta(minutes=15) +# A run whose heartbeat stopped this long ago can be taken over. +STALE_RUN_AFTER = timedelta(minutes=15) +# Per-row failures before the pipeline stops retrying automatically. +MAX_ATTEMPTS = 3 +# Round-level e621 retries before the run is paused. +ROUND_ATTEMPTS = 3 +ROUND_RETRY_SECONDS = 20 + +WORK_STATUSES = (TempUpload.STATUS_PENDING, TempUpload.STATUS_VISUAL_MATCH) +VIDEO_RE = r"\.(mp4|webm)$" + +# A staged upload still needs work when any phase has not run yet. IQDB is +# skipped for videos, which never get iqdb_data, so they must not stay +# "outstanding" forever. +OUTSTANDING_Q = ( + Q(e621_checked_at__isnull=True) + | Q(visual_checked_at__isnull=True) + | (Q(iqdb_data__isnull=True) & ~Q(original_filename__iregex=VIDEO_RE)) +) + +_running_lock = threading.Lock() +_running_users: set[int] = set() + + +class PipelinePaused(Exception): + """A round-level failure that should pause the run instead of failing rows.""" + + +def is_video(filename): + return bool(filename) and filename.lower().endswith((".mp4", ".webm")) + + +def is_finished(temp): + """True when every phase this file needs has run.""" + if temp.status == TempUpload.STATUS_COMPLETED: + return True + if temp.e621_checked_at is None or temp.visual_checked_at is None: + return False + return is_video(temp.original_filename) or temp.iqdb_data is not None + + +def outstanding_queryset(user): + return ( + TempUpload.objects.filter(user=user, status__in=WORK_STATUSES) + .filter(OUTSTANDING_Q) + .filter(attempts__lt=MAX_ATTEMPTS) + ) + + +def count_outstanding(user): + return outstanding_queryset(user).count() + + +def count_failed(user): + return ( + TempUpload.objects.filter(user=user, status__in=WORK_STATUSES) + .filter(attempts__gte=MAX_ATTEMPTS) + .count() + ) + + +def waiting_counts(user): + """How many files are left per phase (phases overlap by design).""" + base = TempUpload.objects.filter( + user=user, status__in=WORK_STATUSES, attempts__lt=MAX_ATTEMPTS + ) + return { + "md5": base.filter(e621_checked_at__isnull=True).count(), + "visual": base.filter(visual_checked_at__isnull=True).count(), + "iqdb": ( + base.filter(iqdb_data__isnull=True) + .exclude(original_filename__iregex=VIDEO_RE) + .count() + ), + } + + +def status_payload(user): + """Cheap state for the shell/upload page to poll.""" + run = UploadRun.objects.filter(user=user).first() + outstanding = count_outstanding(user) + failed = count_failed(user) + status = run.status if run is not None else UploadRun.STATUS_IDLE + total = run.total if run is not None else 0 + processed = run.processed if run is not None else 0 + # A paused run with nothing left to do is not "active" (the user may have + # resolved or discarded the failed rows); failed rows stay visible until + # they are retried or dismissed. + active = ( + outstanding > 0 or failed > 0 or status == UploadRun.STATUS_RUNNING + ) + return { + "status": status, + "active": active, + "phase": run.phase if run is not None else "", + "total": max(total, processed + failed), + "processed": processed, + "matched": run.matched if run is not None else 0, + "failed": max(failed, run.failed if run is not None else 0), + "error": run.error if run is not None else "", + "outstanding": outstanding, + "waiting": waiting_counts(user), + "updated_at": run.updated_at.isoformat() if run is not None else None, + } + + +def reap_stale_claims(): + """Queue rows left claimed by a recycled worker and pause dead runs.""" + cutoff = timezone.now() - STALE_CLAIM_AFTER + released = TempUpload.objects.filter(claimed_at__lt=cutoff).update( + claimed_at=None + ) + paused = UploadRun.objects.filter( + status=UploadRun.STATUS_RUNNING, updated_at__lt=cutoff + ).update( + status=UploadRun.STATUS_PAUSED, + phase="", + error="The worker stopped before finishing. Retry to resume.", + updated_at=timezone.now(), + ) + if released or paused: + logger.info( + "Reaped %s stale upload claims and %s dead upload runs", released, paused + ) + return released, paused + + +def start_pipeline(user): + """Start the pipeline for one user in a daemon thread (idempotent).""" + if not getattr(settings, "UPLOAD_PIPELINE_AUTOSTART", True): + return False + if user is None or not getattr(user, "can_upload", False): + return False + user_id = int(user.pk) + with _running_lock: + if user_id in _running_users: + return False + _running_users.add(user_id) + reap_stale_claims() + thread = threading.Thread(target=_thread_entry, args=(user_id,), daemon=True) + thread.start() + return True + + +def _thread_entry(user_id): + try: + run_pipeline(user_id) + except Exception: # noqa: BLE001 - a thread must never crash the worker + logger.exception("Upload pipeline for user %s crashed", user_id) + finally: + with _running_lock: + _running_users.discard(user_id) + connection.close() + + +def run_pipeline(user_id): + user = get_user_model().objects.filter(pk=user_id).first() + if user is None or not user.can_upload: + return + now = timezone.now() + # Claim the run row so two gunicorn workers cannot own the same queue. + with transaction.atomic(): + run, _ = UploadRun.objects.select_for_update().get_or_create(user=user) + if ( + run.status == UploadRun.STATUS_RUNNING + and run.updated_at is not None + and run.updated_at > now - STALE_RUN_AFTER + ): + # Another worker owns this run. + return + + # Reset the counters when a new queue starts cleanly; otherwise keep + # accumulating so failed rows from an earlier pass stay visible. + live_failed = count_failed(user) + outstanding = count_outstanding(user) + if run.status == UploadRun.STATUS_IDLE and live_failed == 0: + run.total = outstanding + run.processed = 0 + run.matched = 0 + run.failed = 0 + else: + run.total = max( + run.total or 0, run.processed + run.failed + outstanding + ) + run.failed = max(run.failed, live_failed) + run.status = UploadRun.STATUS_RUNNING + run.phase = "" + run.error = "" + run.started_at = now + run.save() + + try: + while True: + rows = claim_round(user) + if not rows: + break + process_round(run, user, rows) + except PipelinePaused as exc: + run.status = UploadRun.STATUS_PAUSED + run.phase = "" + run.error = str(exc) + run.save() + except Exception as exc: # noqa: BLE001 - surface crashes as a run error + logger.exception("Upload pipeline for user %s failed", user_id) + run.status = UploadRun.STATUS_ERROR + run.phase = "" + run.error = f"The upload pipeline stopped: {exc}" + run.save() + else: + run.status = UploadRun.STATUS_IDLE + run.phase = "" + run.save() + + +def claim_round(user, size=CLAIM_SIZE): + """Claim up to ``size`` outstanding rows for this worker.""" + now = timezone.now() + with transaction.atomic(): + rows = list( + TempUpload.objects.select_for_update(skip_locked=True) + .filter(user=user, status__in=WORK_STATUSES) + .filter(OUTSTANDING_Q) + .filter(claimed_at__isnull=True, attempts__lt=MAX_ATTEMPTS) + .order_by("created_at")[:size] + ) + if rows: + TempUpload.objects.filter(pk__in=[row.pk for row in rows]).update( + claimed_at=now + ) + for row in rows: + row.claimed_at = now + return rows + + +def process_round(run, user, rows): + """Run every phase for one claimed round, then release/account the rows.""" + ids = [row.pk for row in rows] + matched = 0 + try: + matched += md5_phase(run, user, rows) + rows = refresh(ids) + visual_phase(run, user, rows) + rows = refresh(ids) + iqdb_phase(run, user, rows) + finally: + finalize_round(run, ids, matched) + + +def refresh(ids): + return list(TempUpload.objects.filter(pk__in=ids)) + + +def _save_run(run, **fields): + for key, value in fields.items(): + setattr(run, key, value) + run.save( + update_fields=[*fields.keys(), "updated_at"] + ) + + +def md5_phase(run, user, rows): + """One e621 MD5 batch query; matches are imported into the library.""" + targets = [row for row in rows if row.e621_checked_at is None] + if not targets: + return 0 + _save_run( + run, + phase=UploadRun.PHASE_MD5, + total=run.processed + run.failed + count_outstanding(user), + ) + posts = _e621_round( + lambda: e621.check_md5_batch(user, [row.md5 for row in targets]) + ) + by_md5 = {} + for post in posts.values(): + file_data = post.get("file") or {} + md5 = str(file_data.get("md5") or "").strip().lower() + if md5: + by_md5[md5] = post + + matched = 0 + now = timezone.now() + for row in targets: + post = by_md5.get(str(row.md5).strip().lower()) + if post is None: + TempUpload.objects.filter(pk=row.pk).update( + e621_checked_at=now, pipeline_error="", updated_at=now + ) + continue + try: + trimmed = services.trim_e621_post(post) + if trimmed is None or not trimmed.get("id"): + raise e621.E621Error("e621 returned an unexpected post payload.") + TempUpload.objects.filter(pk=row.pk).update( + e621_post_id=int(trimmed["id"]), + e621_data=trimmed, + resolution=TempUpload.RESOLUTION_AUTO_MD5, + e621_checked_at=now, + pipeline_error="", + updated_at=now, + ) + row.refresh_from_db() + from .uploads import complete_temp_upload + + complete_temp_upload(row) + matched += 1 + except Exception as exc: # noqa: BLE001 - keep going for other files + logger.exception("Could not auto-import staged upload %s", row.pk) + record_failure(row, f"Could not finish the upload: {exc}") + return matched + + +def visual_phase(run, user, rows): + """Compare each row's perceptual hashes against the library once.""" + from .uploads import build_hash_index, match_hashes + + targets = [ + row + for row in rows + if row.visual_checked_at is None + and row.status in WORK_STATUSES + and row.file + ] + if not targets: + return + _save_run(run, phase=UploadRun.PHASE_VISUAL) + index = build_hash_index() + now = timezone.now() + for row in targets: + try: + hashes = services.compute_visual_hashes(row.file.path) + if not hashes: + TempUpload.objects.filter(pk=row.pk).update( + visual_checked_at=now, pipeline_error="", updated_at=now + ) + continue + matches = match_hashes(hashes, index, user=user) + update = { + "visual_matches": matches, + "visual_checked_at": now, + "pipeline_error": "", + "updated_at": now, + } + if matches and row.status == TempUpload.STATUS_PENDING: + update["status"] = TempUpload.STATUS_VISUAL_MATCH + TempUpload.objects.filter(pk=row.pk).update(**update) + except Exception as exc: # noqa: BLE001 - keep going for other files + logger.exception("Visual similarity failed for %s", row.pk) + record_failure(row, f"Visual similarity failed: {exc}") + + +def iqdb_phase(run, user, rows): + """Reverse-image search every unresolved image, one e621 query each.""" + targets = [ + row + for row in rows + if row.iqdb_data is None + and row.status in WORK_STATUSES + and row.file + and not is_video(row.original_filename) + ] + if not targets: + return + _save_run(run, phase=UploadRun.PHASE_IQDB) + heartbeat_at = time.monotonic() + for row in targets: + # Keep the run row fresh: a 75-file IQDB round takes minutes and must + # not look like a dead worker to another request. + if time.monotonic() - heartbeat_at > 30: + UploadRun.objects.filter(pk=run.pk).update(updated_at=timezone.now()) + heartbeat_at = time.monotonic() + try: + raw = e621.iqdb_search(user, row.file.path) + results = normalize_iqdb_results(user, raw) + except e621.E621AuthError as exc: + record_failure(row, str(exc)) + raise PipelinePaused( + "e621 rejected the credentials — fix them in Account and retry." + ) from exc + except e621.E621RateLimited as exc: + record_failure(row, str(exc)) + raise PipelinePaused( + "e621 is throttling IQDB right now; the queue will resume." + ) from exc + except Exception as exc: # noqa: BLE001 - keep going for other files + logger.exception("IQDB search failed for %s", row.pk) + record_failure(row, f"IQDB search failed: {exc}") + continue + now = timezone.now() + update = { + "iqdb_data": results, + "pipeline_error": "", + "updated_at": now, + } + if results and row.status == TempUpload.STATUS_PENDING: + update["status"] = TempUpload.STATUS_VISUAL_MATCH + TempUpload.objects.filter(pk=row.pk).update(**update) + + +def record_failure(row, message): + """Count one failed attempt against a row and queue it for a retry.""" + TempUpload.objects.filter(pk=row.pk).update( + attempts=F("attempts") + 1, + pipeline_error=str(message)[:2000], + claimed_at=None, + updated_at=timezone.now(), + ) + + +def finalize_round(run, ids, matched): + """Account finished/failed rows and release the rest of the claims.""" + rows = refresh(ids) + finished = {row.pk for row in rows if is_finished(row)} + failed = {row.pk for row in rows if row.attempts >= MAX_ATTEMPTS} + # Release every claim: finished rows must not keep looking "processing" + # to the board, and unfinished rows are re-queued for the next run. + release = [row.pk for row in rows if row.claimed_at is not None] + if release: + TempUpload.objects.filter(pk__in=release).update(claimed_at=None) + _save_run( + run, + processed=run.processed + len(finished), + failed=run.failed + len(failed - finished), + matched=run.matched + matched, + ) + + +def _e621_round(task): + """Run a round-level e621 call, retrying through rate limits.""" + last_error = None + for attempt in range(ROUND_ATTEMPTS): + try: + return task() + except e621.E621AuthError: + raise + except (e621.E621RateLimited, e621.E621Error) as exc: + last_error = exc + if attempt + 1 >= ROUND_ATTEMPTS: + break + delay = ROUND_RETRY_SECONDS * (attempt + 1) + logger.info("e621 round failed (%s); retrying in %ss", exc, delay) + time.sleep(delay) + raise PipelinePaused( + f"e621 is not answering right now ({last_error}); the queue will resume." + ) from last_error + + +def flatten_tag_preview(tags, limit=8): + """First few tag names from a modern post payload, like the SPA shows.""" + if not isinstance(tags, dict): + return [] + out = [] + for values in tags.values(): + if not isinstance(values, list): + continue + for tag in values: + if isinstance(tag, str) and tag not in out: + out.append(tag) + if len(out) >= limit: + return out + return out + + +def _legacy_iqdb_post(entry): + """Unwrap the post payload embedded in a legacy IQDB match.""" + post = entry.get("post") + if not isinstance(post, dict): + return {} + inner = post.get("posts") + return inner if isinstance(inner, dict) else post + + +def normalize_iqdb_results(user, raw_results): + """Shape legacy IQDB matches like the SPA's E621IqdbCandidate entries. + + The IQDB payload carries little post data, so candidates are enriched + with one batched ``id:`` lookup before they are stored. + """ + candidates = [] + for entry in (raw_results or [])[:10]: + if not isinstance(entry, dict): + continue + post = _legacy_iqdb_post(entry) + post_id = entry.get("post_id") + if not isinstance(post_id, int): + post_id = post.get("id") + score = entry.get("score") + candidates.append( + { + "post_id": post_id if isinstance(post_id, int) else None, + "score": float(score) if isinstance(score, (int, float)) else None, + "preview_url": None, + "rating": ( + post.get("rating") if isinstance(post.get("rating"), str) else None + ), + "md5": post.get("md5") if isinstance(post.get("md5"), str) else None, + "score_total": ( + post.get("score") if isinstance(post.get("score"), int) else None + ), + "fav_count": ( + post.get("fav_count") + if isinstance(post.get("fav_count"), int) + else None + ), + "width": ( + post.get("image_width") + if isinstance(post.get("image_width"), int) + else None + ), + "height": ( + post.get("image_height") + if isinstance(post.get("image_height"), int) + else None + ), + "tags_preview": [], + } + ) + + ids = [entry["post_id"] for entry in candidates if entry["post_id"]] + if not ids: + return services.sanitize_iqdb_results(candidates) + + try: + posts = e621.fetch_posts_by_ids(user, ids) + except e621.E621Error as exc: + # Candidates without enrichment still show up; keep them. + logger.info("Could not enrich IQDB candidates: %s", exc) + posts = [] + by_id = {post.get("id"): post for post in posts} + + for entry in candidates: + post = by_id.get(entry["post_id"]) + if not isinstance(post, dict): + continue + file_data = post.get("file") or {} + preview = post.get("preview") or {} + score = post.get("score") or {} + entry["preview_url"] = preview.get("url") or entry["preview_url"] + entry["rating"] = post.get("rating") or entry["rating"] + entry["md5"] = file_data.get("md5") or entry["md5"] + if isinstance(score, dict): + entry["score_total"] = score.get("total") + entry["fav_count"] = post.get("fav_count") + entry["width"] = file_data.get("width") + entry["height"] = file_data.get("height") + entry["tags_preview"] = flatten_tag_preview(post.get("tags")) + + return services.sanitize_iqdb_results(candidates) diff --git a/backend/apps/library/uploads.py b/backend/apps/library/uploads.py index 09ef3f1..25fa899 100644 --- a/backend/apps/library/uploads.py +++ b/backend/apps/library/uploads.py @@ -29,27 +29,32 @@ from rest_framework.response import Response from . import services from .models import MediaItem, TempUpload from .permissions import CanUpload -from .serializers import TempUploadSerializer -from .tools import HASH_FIELDS, hashes_similarity +from .serializers import TempUploadListSerializer, TempUploadSerializer +from .tools import HASH_FIELDS, hashed_items, hashes_similarity logger = logging.getLogger(__name__) -def find_library_matches(path, limit=10, user=None, request=None): - """Library items visually similar to a staged file.""" - hashes = services.compute_visual_hashes(path) - if not hashes: - return [] +def build_hash_index(): + """Library hash mappings for similarity scans, loaded once per batch. + + Only items that actually carry perceptual hashes are included; the old + per-file scan walked every row (including videos and unchecked items). + """ + algorithms = list(HASH_FIELDS) + return [ + (item, {field: getattr(item, field, "") for field in algorithms}) + for item in hashed_items(algorithms) + ] + + +def match_hashes(hashes, index, limit=10, user=None, request=None): + """Library items whose perceptual hashes are close to ``hashes``.""" algorithms = list(HASH_FIELDS) threshold = settings.VISUAL_MATCH_THRESHOLD matches = [] - for item in MediaItem.objects.prefetch_related("locations"): - similarity = hashes_similarity( - hashes, - {field: getattr(item, field, "") for field in algorithms}, - algorithms, - threshold, - ) + for item, item_hashes in index: + similarity = hashes_similarity(hashes, item_hashes, algorithms, threshold) if similarity is None: continue location = item.locations.first() @@ -67,6 +72,20 @@ def find_library_matches(path, limit=10, user=None, request=None): return matches[:limit] +def find_library_matches(path, limit=10, user=None, request=None, index=None): + """Library items visually similar to a staged file. + + Pass a prebuilt ``index`` (see build_hash_index) to reuse it across a + whole batch instead of rescanning the library per file. + """ + hashes = services.compute_visual_hashes(path) + if not hashes: + return [] + if index is None: + index = build_hash_index() + return match_hashes(hashes, index, limit=limit, user=user, request=request) + + def complete_temp_upload(temp, download_url=None): """Index the upload into the library. @@ -156,12 +175,20 @@ class TempUploadViewSet( pagination_class = None http_method_names = ["get", "post", "delete", "head", "options"] + def get_serializer_class(self): + # The board polls the list, so its payload stays small; the metadata + # modal fetches the full row from the detail endpoint. + if self.action == "list": + return TempUploadListSerializer + return TempUploadSerializer + def get_queryset(self): - queryset = TempUpload.objects.select_related("library_item") - user = self.request.user - if not user.is_app_staff: - queryset = queryset.filter(user=user) - return queryset + # Staged uploads are private: everyone, staff included, only sees + # their own board. (The file action still lets staff read bytes by id + # for support purposes.) + return TempUpload.objects.select_related("library_item").filter( + user=self.request.user + ) def create(self, request): upload = request.FILES.get("file") @@ -190,10 +217,14 @@ class TempUploadViewSet( temp.resolution = TempUpload.RESOLUTION_DUPLICATE temp.library_item = existing temp.file.delete(save=False) - # Visual similarity is deliberately a separate phase (the - # /visual-match action) so a large batch uploads at full speed and - # the board runs MD5 -> visual -> IQDB over the whole batch. + # Visual similarity and IQDB run in the background pipeline so a large + # batch uploads at full speed and the work survives the browser. temp.save() + # Kick the server-side pipeline; staging no longer waits on e621 and + # the work continues even if the browser navigates away. + from .upload_pipeline import start_pipeline + + start_pipeline(request.user) return Response( self.get_serializer(temp).data, status=status.HTTP_201_CREATED ) @@ -221,6 +252,135 @@ class TempUploadViewSet( temp.save(update_fields=["visual_matches", "status", "updated_at"]) return Response(self.get_serializer(temp).data) + @action(detail=False, methods=["get"]) + def status(self, request): + """Cheap pipeline state for the shell indicator and the upload page.""" + from .upload_pipeline import status_payload + + return Response(status_payload(request.user)) + + @action(detail=False, methods=["post"]) + def process(self, request): + """Start (or resume) the pipeline for the caller's staged uploads. + + Idempotent: the client calls this after staging files, on page load + and when a paused run should be retried. + """ + from .upload_pipeline import start_pipeline, status_payload + + start_pipeline(request.user) + return Response(status_payload(request.user)) + + @action(detail=True, methods=["post"]) + def retry(self, request, pk=None): + """Queue one staged upload for another pipeline pass.""" + from .upload_pipeline import start_pipeline, status_payload + + temp = self.get_object() + if temp.status == TempUpload.STATUS_COMPLETED: + return Response( + {"detail": "This upload is already in the library."}, + status=status.HTTP_400_BAD_REQUEST, + ) + if not temp.file: + return Response( + {"detail": "The staged file is missing."}, + status=status.HTTP_400_BAD_REQUEST, + ) + update = { + "claimed_at": None, + "attempts": 0, + "pipeline_error": "", + "updated_at": timezone.now(), + } + # An explicit phase re-runs that one check even if it already ran. + phase = str(request.data.get("phase") or "").strip() + if phase == "md5": + update["e621_checked_at"] = None + elif phase == "visual": + update["visual_matches"] = None + update["visual_checked_at"] = None + elif phase == "iqdb": + update["iqdb_data"] = None + if temp.status == TempUpload.STATUS_ERROR: + # A failed import has to go through the MD5 phase again so the + # completion is retried; other errors only re-run missing phases. + update["status"] = TempUpload.STATUS_PENDING + update["e621_checked_at"] = None + elif temp.status not in ( + TempUpload.STATUS_PENDING, + TempUpload.STATUS_VISUAL_MATCH, + ): + update["status"] = TempUpload.STATUS_PENDING + TempUpload.objects.filter(pk=temp.pk).update(**update) + start_pipeline(request.user) + return Response(status_payload(request.user)) + + @action(detail=False, methods=["post"], url_path="retry-all") + def retry_all(self, request): + """Queue every retryable staged upload for another pipeline pass.""" + from .upload_pipeline import start_pipeline, status_payload + + now = timezone.now() + retryable = self.get_queryset().filter( + status__in=[ + TempUpload.STATUS_PENDING, + TempUpload.STATUS_VISUAL_MATCH, + TempUpload.STATUS_ERROR, + ] + ) + retryable.exclude(file="").update( + claimed_at=None, + attempts=0, + pipeline_error="", + updated_at=now, + ) + retryable.exclude(file="").filter(status=TempUpload.STATUS_ERROR).update( + status=TempUpload.STATUS_PENDING, + e621_checked_at=None, + ) + start_pipeline(request.user) + return Response(status_payload(request.user)) + + @action(detail=False, methods=["post"], url_path="discard-bulk") + def discard_bulk(self, request): + """Discard many staged uploads in one request (the board's "all").""" + ids = request.data.get("temp_ids") + if not isinstance(ids, list) or not ids: + return Response( + {"detail": "temp_ids must be a non-empty list."}, + status=status.HTTP_400_BAD_REQUEST, + ) + if len(ids) > 1000: + return Response( + {"detail": "Too many ids in one request (max 1000)."}, + status=status.HTTP_400_BAD_REQUEST, + ) + values = list(dict.fromkeys(str(value) for value in ids)) + try: + queryset = self.get_queryset().filter(pk__in=values) + except (ValidationError, ValueError): + return Response( + {"detail": "One or more temp_ids are not valid upload ids."}, + status=status.HTTP_400_BAD_REQUEST, + ) + by_id = {str(temp.pk): temp for temp in queryset} + + discarded: list[str] = [] + errors: list[dict[str, str]] = [] + for value in values: + temp = by_id.get(value) + if temp is None: + errors.append({"temp_id": value, "error": "not found"}) + continue + try: + self.perform_destroy(temp) + discarded.append(value) + except Exception as exc: # noqa: BLE001 - report per-file failures + logger.exception("Could not discard staged upload %s", value) + errors.append({"temp_id": value, "error": str(exc)}) + return Response({"discarded": discarded, "errors": errors}) + @action(detail=True, methods=["get", "head"], permission_classes=[AllowAny]) def file(self, request, pk=None): """Serve the staged file; accepts a signed URL for media tags.""" diff --git a/backend/config/settings.py b/backend/config/settings.py index 3e4d74b..8188541 100644 --- a/backend/config/settings.py +++ b/backend/config/settings.py @@ -234,6 +234,12 @@ GUEST_BLACKLIST_TTL = int(os.getenv("GUEST_BLACKLIST_TTL", "3600")) # Similarity threshold for flagging staged uploads that match library items. VISUAL_MATCH_THRESHOLD = float(os.getenv("VISUAL_MATCH_THRESHOLD", "0.9")) +# Start the staged-upload pipeline when a file is staged (daemon thread in the +# worker). Tests turn this off and drive the pipeline synchronously. +UPLOAD_PIPELINE_AUTOSTART = os.getenv( + "UPLOAD_PIPELINE_AUTOSTART", "true" +).strip().lower() not in {"0", "false", "no", "off"} + # Ephemeral similarity-check uploads are deleted after this many minutes # (and always on startup). SIMILARITY_TTL_MINUTES = int(os.getenv("SIMILARITY_TTL_MINUTES", "30")) diff --git a/frontend/src/components/AppShell.tsx b/frontend/src/components/AppShell.tsx index 4caee20..830457f 100644 --- a/frontend/src/components/AppShell.tsx +++ b/frontend/src/components/AppShell.tsx @@ -15,6 +15,7 @@ import { ConfirmDialog } from "@/components/ConfirmDialog"; import { StatusFooter } from "@/components/StatusFooter"; import { StatusPill } from "@/components/StatusPill"; import { Toasts } from "@/components/Toasts"; +import { UploadIndicator } from "@/components/UploadIndicator"; import { api } from "@/lib/api"; import { hasBackend } from "@/lib/backend"; import { cn } from "@/lib/cn"; @@ -153,6 +154,7 @@ export function AppShell() {
+ {phases.map(([label, done], index) => ( + + {index > 0 ? " · " : ""} + + {label} {done ? "✓" : "…"} + + + ))} +
+ ); +} + function TempCard({ temp, - checking, onOpen, - onCheck, + onRetry, onDismiss, }: { temp: TempUpload; - checking: boolean; onOpen: () => void; - onCheck: () => void; + onRetry: () => void; onDismiss: () => void; }) { const isVideo = /\.(mp4|webm)$/i.test(temp.original_filename); const preview = apiUrl(temp.preview_url ?? temp.file_url ?? ""); + const similarCount = temp.similar_count ?? 0; return (- Something went wrong while indexing this file. -
+ {temp.pipeline_error ? ( +{temp.pipeline_error}
) : null} {temp.status === "pending" || temp.status === "visual_match" ? ( @@ -244,10 +259,14 @@ function TempCard({ )} @@ -270,17 +289,18 @@ function TempCard({ ); } -function MetadataModal({ +/** + * Full metadata editor for one staged upload. + * + * The board only carries compact rows, so this fetches the full record (post + * payload, IQDB candidates, visual matches) and keeps polling while the + * background pipeline is still working on it. + */ +function MetadataForm({ temp, - checking, - checkError, - onCheck, onClose, }: { temp: TempUpload; - checking: boolean; - checkError: string | null; - onCheck: () => void; onClose: () => void; }) { const queryClient = useQueryClient(); @@ -289,21 +309,30 @@ function MetadataModal({ const [postId, setPostId] = useState(""); const [selected, setSelected] = useState+ {temp.original_filename} +
++ Already in your library +
++ This upload looks like it is already in the library — you + can discard it below. +
+ ++ IQDB candidates +
+ {isVideo ? null : ( + + )} ++ IQDB works on images only. +
+ ) : temp.processing && !(iqdbData?.length ?? 0) ? ( +
+
+ Post #{selected.post_id} +
++ {selected.rating + ? (RATING_LABELS[selected.rating] ?? + selected.rating) + : "Unknown rating"} + {selected.score_total !== null && + selected.score_total !== undefined + ? ` · ▲ ${selected.score_total}` + : ""} + {selected.fav_count !== null && + selected.fav_count !== undefined + ? ` · ${selected.fav_count} favs` + : ""} + {selected.width && selected.height + ? ` · ${selected.width}×${selected.height}` + : ""} +
+ {selected.tags_preview && + selected.tags_preview.length > 0 ? ( ++ Linking downloads the post's file into the library + and drops this staged upload. +
++ Checked against e621 IQDB — no match found. +
+ ) : ( ++ Not checked yet. The background pipeline checks uploads + automatically. +
+ )} +{error}
: null} + {temp.pipeline_error ? ( ++ Pipeline error: {temp.pipeline_error} +
+ ) : null} + {busy ? ( +
+
- {temp.original_filename} -
-- Already in your library -
-- This upload looks like it is already in the library — - you can discard it below. -
- -- IQDB candidates -
- {isVideo || !credentials?.configured ? null : ( - - )} -- IQDB works on images only. -
- ) : !credentials?.configured ? ( -- Configure e621 credentials in Account to use IQDB. -
- ) : checking && !(temp.iqdb_data?.length ?? 0) ? ( -
-
- IQDB check failed: {checkError} -
- ) : temp.iqdb_data && temp.iqdb_data.length > 0 ? ( - <> -- Post #{selected.post_id} -
-- {selected.rating - ? (RATING_LABELS[selected.rating] ?? - selected.rating) - : "Unknown rating"} - {selected.score_total !== null && - selected.score_total !== undefined - ? ` · ▲ ${selected.score_total}` - : ""} - {selected.fav_count !== null && - selected.fav_count !== undefined - ? ` · ${selected.fav_count} favs` - : ""} - {selected.width && selected.height - ? ` · ${selected.width}×${selected.height}` - : ""} -
- {selected.tags_preview && - selected.tags_preview.length > 0 ? ( -- Linking downloads the post's file into the library - and drops this staged upload. -
-- Checked against e621 IQDB — no match found. -
- ) : ( -- Not checked yet. Uploads are checked automatically after - they finish. -
- )} -{error}
- ) : null} - {busy ? ( -
-
+ Could not load this upload's details. +
+ ) : ( +
+
- Files are staged first, auto-matched by MD5, checked against IQDB, - then moved into the library once resolved. + Files are staged first, then a background worker matches them by + MD5, checks visual similarity and queries IQDB — you can leave this + page while it runs.
IQDB checks paused: {iqdbError}
+ {status && (status.status === "paused" || status.status === "error") ? ( ++ {status.status === "paused" + ? "Upload processing paused" + : "Upload processing stopped"} + : {status.error} +