Upload updates

This commit is contained in:
2026-09-21 09:01:01 -05:00
parent 98025e9e6d
commit 474403ffe2
18 changed files with 2539 additions and 796 deletions
+589
View File
@@ -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)