Upload updates
This commit is contained in:
+12
-9
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
+257
-34
@@ -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}",
|
||||
return _request(
|
||||
user,
|
||||
"GET",
|
||||
path,
|
||||
params=params,
|
||||
auth=(
|
||||
(user.e621_username, user.e621_api_key_plain)
|
||||
if configured
|
||||
else None
|
||||
),
|
||||
headers={"User-Agent": settings.USER_AGENT},
|
||||
timeout=timeout,
|
||||
require_auth=require_auth,
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
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 []
|
||||
|
||||
@@ -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"]:
|
||||
|
||||
@@ -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)."))
|
||||
@@ -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
|
||||
),
|
||||
]
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
+181
-21
@@ -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")
|
||||
# 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
|
||||
if not user.is_app_staff:
|
||||
queryset = queryset.filter(user=user)
|
||||
return queryset
|
||||
)
|
||||
|
||||
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."""
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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() {
|
||||
<div className="ml-auto flex items-center gap-3">
|
||||
{backend ? (
|
||||
<>
|
||||
<UploadIndicator />
|
||||
<StatusPill status={status} />
|
||||
{user ? (
|
||||
<>
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
import { Link } from "react-router-dom";
|
||||
|
||||
import { cn } from "@/lib/cn";
|
||||
import { useUploadStatus } from "@/lib/uploadStatus";
|
||||
|
||||
const PHASE_LABELS: Record<string, string> = {
|
||||
md5: "e621 MD5",
|
||||
visual: "Visual similarity",
|
||||
iqdb: "IQDB",
|
||||
};
|
||||
|
||||
/**
|
||||
* Shell-wide progress for the server-side upload pipeline.
|
||||
*
|
||||
* Staged uploads keep processing after the Upload page is closed, so this
|
||||
* pill is the "leave and keep an eye on it" view: it lives in the header on
|
||||
* every page and links back to the board.
|
||||
*/
|
||||
export function UploadIndicator() {
|
||||
const { data: status } = useUploadStatus();
|
||||
if (!status || !status.active) return null;
|
||||
|
||||
const total = Math.max(status.total, status.processed + status.failed);
|
||||
const done = status.processed + status.failed;
|
||||
const percent = total > 0 ? Math.min(100, Math.round((done / total) * 100)) : 0;
|
||||
const broken = status.status === "paused" || status.status === "error";
|
||||
const label =
|
||||
status.status === "paused"
|
||||
? "uploads paused"
|
||||
: status.status === "error"
|
||||
? "uploads failed"
|
||||
: (PHASE_LABELS[status.phase] ?? "processing uploads");
|
||||
const count = total > 0 ? `${done}/${total}` : `${status.outstanding} left`;
|
||||
|
||||
return (
|
||||
<Link
|
||||
to="/upload"
|
||||
title={status.error || "Staged uploads are being processed in the background"}
|
||||
className={cn(
|
||||
"flex items-center gap-2 rounded-md border px-2 py-1 font-mono text-[10px] transition",
|
||||
broken
|
||||
? "border-ctp-red/40 text-ctp-red hover:bg-ctp-red/10"
|
||||
: "border-ctp-surface1 text-ctp-subtext0 hover:bg-ctp-surface0 hover:text-ctp-text",
|
||||
)}
|
||||
>
|
||||
<span className="hidden sm:inline">{label}</span>
|
||||
<span>{count}</span>
|
||||
<span className="h-1 w-12 overflow-hidden rounded-full bg-ctp-surface0">
|
||||
<span
|
||||
className={cn("block h-full", broken ? "bg-ctp-red" : "bg-ctp-teal")}
|
||||
style={{ width: `${percent}%` }}
|
||||
/>
|
||||
</span>
|
||||
</Link>
|
||||
);
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -355,12 +355,17 @@ export interface TempUpload {
|
||||
status: "pending" | "visual_match" | "completed" | "error";
|
||||
resolution: "" | "auto_md5" | "duplicate" | "linked" | "custom";
|
||||
e621_post_id: number | null;
|
||||
e621_data: E621StoredPost | null;
|
||||
/** Full post payload: only present on the detail endpoint. */
|
||||
e621_data?: E621StoredPost | null;
|
||||
custom_rating: Rating;
|
||||
custom_tags: string[];
|
||||
custom_notes: string;
|
||||
iqdb_data: E621IqdbCandidate[] | null;
|
||||
visual_matches:
|
||||
/** Only present on the detail endpoint. */
|
||||
custom_tags?: string[];
|
||||
/** Only present on the detail endpoint. */
|
||||
custom_notes?: string;
|
||||
/** Only present on the detail endpoint. */
|
||||
iqdb_data?: E621IqdbCandidate[] | null;
|
||||
/** Only present on the detail endpoint. */
|
||||
visual_matches?:
|
||||
| {
|
||||
j_id: string;
|
||||
filename: string;
|
||||
@@ -371,10 +376,33 @@ export interface TempUpload {
|
||||
library_j_id: string | null;
|
||||
file_url: string | null;
|
||||
preview_url: string | null;
|
||||
/** Background pipeline state (see the UploadRun status endpoint). */
|
||||
pipeline_error: string;
|
||||
attempts: number;
|
||||
md5_checked?: boolean;
|
||||
visual_checked?: boolean;
|
||||
iqdb_checked?: boolean;
|
||||
processing?: boolean;
|
||||
similar_count?: number;
|
||||
created_at: string;
|
||||
updated_at: string;
|
||||
}
|
||||
|
||||
/** Progress of the server-side staged-upload pipeline. */
|
||||
export interface UploadStatus {
|
||||
status: "idle" | "running" | "paused" | "error";
|
||||
active: boolean;
|
||||
phase: "" | "md5" | "visual" | "iqdb";
|
||||
total: number;
|
||||
processed: number;
|
||||
matched: number;
|
||||
failed: number;
|
||||
error: string;
|
||||
outstanding: number;
|
||||
waiting: { md5: number; visual: number; iqdb: number };
|
||||
updated_at: string | null;
|
||||
}
|
||||
|
||||
export interface StatJob {
|
||||
kind: "download" | "match";
|
||||
id: string;
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
import { useQuery } from "@tanstack/react-query";
|
||||
|
||||
import { api } from "@/lib/api";
|
||||
import { hasBackend } from "@/lib/backend";
|
||||
import type { UploadStatus } from "@/lib/types";
|
||||
import { useAuth } from "@/store/auth";
|
||||
|
||||
export const uploadStatusQueryKey = ["upload-status"] as const;
|
||||
|
||||
/**
|
||||
* Poll the server-side upload pipeline while it has work and stop once it is
|
||||
* idle. Shared by the Upload page and the shell indicator so both show the
|
||||
* same state without extra requests.
|
||||
*/
|
||||
export function useUploadStatus(enabled = true) {
|
||||
const user = useAuth((state) => state.user);
|
||||
return useQuery({
|
||||
queryKey: uploadStatusQueryKey,
|
||||
queryFn: () => api<UploadStatus>("/api/uploads/status/"),
|
||||
enabled: enabled && hasBackend() && Boolean(user?.can_upload),
|
||||
// TanStack takes the smallest interval across the observers, so the
|
||||
// Upload page and the shell pill can safely share this query.
|
||||
refetchInterval: (query) => (query.state.data?.active ? 2_000 : false),
|
||||
});
|
||||
}
|
||||
|
||||
/** Ask the server to (re)start processing the caller's staged uploads. */
|
||||
export function postProcessUploads(): Promise<UploadStatus> {
|
||||
return api<UploadStatus>("/api/uploads/process/", { method: "POST" });
|
||||
}
|
||||
Reference in New Issue
Block a user