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