590 lines
20 KiB
Python
590 lines
20 KiB
Python
"""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)
|