Uploads were doing md5 + local visual matching inside the upload request (backend create) while the frontend later ran its own e621 MD5 pass, so the pipeline looked interleaved per file. Now every step is a phase applied to the whole batch in order: 1. upload (fast: md5 + exact-duplicate check only), 2. e621 MD5 lookup, 75 md5: metatags per posts.json request, 3. local visual similarity, one file at a time via the new POST /api/uploads/<id>/visual-match/ action, 4. IQDB through the existing serial queue. The board shows the active phase with its own progress bar (e621 MD5 in peach, visual in lavender, IQDB in teal) and every step updates the staged list as it lands. Verified from a headless run: one batched posts.json request for 10 files, then 10 visual-match calls, then IQDB.
348 lines
13 KiB
Python
348 lines
13 KiB
Python
"""Staged uploads: listing, phases (MD5/visual/IQDB) and the bulk tool."""
|
|
|
|
import base64
|
|
import hashlib
|
|
import io
|
|
import json
|
|
import shutil
|
|
import tempfile
|
|
from pathlib import Path
|
|
|
|
from django.contrib.auth import get_user_model
|
|
from django.core.files.uploadedfile import SimpleUploadedFile
|
|
from django.test import Client, TestCase, override_settings
|
|
|
|
from PIL import Image
|
|
|
|
from rest_framework.authtoken.models import Token
|
|
|
|
from apps.library.models import MediaItem, TempUpload
|
|
|
|
User = get_user_model()
|
|
|
|
# 1x1 transparent PNG so indexing/hashing has a real image to chew on.
|
|
TINY_PNG = base64.b64decode(
|
|
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
|
|
)
|
|
|
|
|
|
def similar_png(color=(200, 30, 40)):
|
|
"""A 1x1 PNG with the same visual hashes but a different MD5."""
|
|
buffer = io.BytesIO()
|
|
Image.new("RGB", (1, 1), color).save(buffer, format="PNG")
|
|
return buffer.getvalue()
|
|
|
|
|
|
def jpost(client, path, body=None):
|
|
return client.post(path, data=json.dumps(body or {}), content_type="application/json")
|
|
|
|
|
|
class TempUploadListTests(TestCase):
|
|
def setUp(self):
|
|
self.uploader = User.objects.create_user(
|
|
username="upload-user", password="upload-pass-123456"
|
|
)
|
|
self.uploader.role = "uploader"
|
|
self.uploader.save(update_fields=["role"])
|
|
self.other = User.objects.create_user(
|
|
username="upload-other", password="upload-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 test_list_is_not_paginated_past_48(self):
|
|
TempUpload.objects.bulk_create(
|
|
[
|
|
TempUpload(
|
|
user=self.uploader,
|
|
original_filename=f"file_{index:03d}.png",
|
|
md5=hashlib.md5(f"file-{index}".encode()).hexdigest(),
|
|
size=index,
|
|
)
|
|
for index in range(69)
|
|
]
|
|
)
|
|
response = self.api_client(self.uploader).get("/api/uploads/")
|
|
self.assertEqual(response.status_code, 200)
|
|
data = response.json()
|
|
self.assertIsInstance(data, list, "the list must not be paginated")
|
|
self.assertEqual(len(data), 69)
|
|
|
|
def test_list_only_contains_the_callers_uploads(self):
|
|
TempUpload.objects.create(
|
|
user=self.uploader,
|
|
original_filename="mine.png",
|
|
md5=hashlib.md5(b"mine").hexdigest(),
|
|
size=1,
|
|
)
|
|
TempUpload.objects.create(
|
|
user=self.other,
|
|
original_filename="theirs.png",
|
|
md5=hashlib.md5(b"theirs").hexdigest(),
|
|
size=1,
|
|
)
|
|
mine = self.api_client(self.uploader).get("/api/uploads/").json()
|
|
self.assertEqual([row["original_filename"] for row in mine], ["mine.png"])
|
|
|
|
def test_anonymous_cannot_list(self):
|
|
self.assertEqual(Client().get("/api/uploads/").status_code, 401)
|
|
|
|
|
|
class StagedUploadWorkflowTests(TestCase):
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
super().setUpClass()
|
|
cls._tmp = tempfile.mkdtemp(prefix="j621-bulk-")
|
|
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="bulk-uploader", password="bulk-pass-123456"
|
|
)
|
|
self.uploader.role = "uploader"
|
|
self.uploader.save(update_fields=["role"])
|
|
self.other = User.objects.create_user(
|
|
username="bulk-other", password="bulk-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 make_temp(self, user, label, status=TempUpload.STATUS_PENDING):
|
|
return TempUpload.objects.create(
|
|
user=user,
|
|
file=SimpleUploadedFile(f"{label}.png", TINY_PNG, content_type="image/png"),
|
|
original_filename=f"{label}.png",
|
|
md5=hashlib.md5(label.encode()).hexdigest(),
|
|
size=len(TINY_PNG),
|
|
status=status,
|
|
)
|
|
|
|
def resolve_bulk(self, client, ids, rating):
|
|
return jpost(
|
|
client,
|
|
"/api/uploads/resolve-bulk/",
|
|
{"temp_ids": [str(value) for value in ids], "rating": rating},
|
|
)
|
|
|
|
def test_resolves_selected_uploads_with_the_rating(self):
|
|
client = self.api_client(self.uploader)
|
|
first = self.make_temp(self.uploader, "one")
|
|
second = self.make_temp(self.uploader, "two")
|
|
untouched = self.make_temp(self.uploader, "three")
|
|
|
|
response = self.resolve_bulk(client, [first.id, second.id], "q")
|
|
self.assertEqual(response.status_code, 200)
|
|
body = response.json()
|
|
self.assertEqual(len(body["resolved"]), 2)
|
|
self.assertEqual(body["errors"], [])
|
|
|
|
for temp in (first, second):
|
|
temp.refresh_from_db()
|
|
self.assertEqual(temp.status, TempUpload.STATUS_COMPLETED)
|
|
self.assertIsNotNone(temp.library_item_id)
|
|
self.assertEqual(temp.library_item.rating, "q")
|
|
self.assertEqual(temp.library_item.uploaded_by_id, self.uploader.id)
|
|
untouched.refresh_from_db()
|
|
self.assertEqual(untouched.status, TempUpload.STATUS_PENDING)
|
|
|
|
def test_rejects_bad_input(self):
|
|
client = self.api_client(self.uploader)
|
|
temp = self.make_temp(self.uploader, "input")
|
|
self.assertEqual(self.resolve_bulk(client, [], "s").status_code, 400)
|
|
self.assertEqual(self.resolve_bulk(client, [temp.id], "").status_code, 400)
|
|
self.assertEqual(self.resolve_bulk(client, [temp.id], "x").status_code, 400)
|
|
self.assertEqual(
|
|
jpost(
|
|
client,
|
|
"/api/uploads/resolve-bulk/",
|
|
{"temp_ids": ["not-a-uuid"], "rating": "s"},
|
|
).status_code,
|
|
400,
|
|
)
|
|
|
|
def test_other_users_uploads_are_left_alone(self):
|
|
client = self.api_client(self.uploader)
|
|
theirs = self.make_temp(self.other, "theirs")
|
|
response = self.resolve_bulk(client, [theirs.id], "s")
|
|
self.assertEqual(response.status_code, 200)
|
|
body = response.json()
|
|
self.assertEqual(body["resolved"], [])
|
|
self.assertEqual(body["errors"][0]["error"], "not found")
|
|
theirs.refresh_from_db()
|
|
self.assertEqual(theirs.status, TempUpload.STATUS_PENDING)
|
|
|
|
def test_completed_uploads_report_an_error_but_others_resolve(self):
|
|
client = self.api_client(self.uploader)
|
|
done = self.make_temp(self.uploader, "done", status=TempUpload.STATUS_COMPLETED)
|
|
pending = self.make_temp(self.uploader, "pending")
|
|
response = self.resolve_bulk(client, [done.id, pending.id], "e")
|
|
body = response.json()
|
|
self.assertEqual(len(body["resolved"]), 1)
|
|
self.assertEqual(body["errors"][0]["error"], "already in the library")
|
|
self.assertEqual(MediaItem.objects.count(), 1)
|
|
|
|
def upload_via_api(self, client, label="phase.png"):
|
|
return client.post(
|
|
"/api/uploads/",
|
|
{
|
|
"file": SimpleUploadedFile(
|
|
label, similar_png(), content_type="image/png"
|
|
)
|
|
},
|
|
)
|
|
|
|
def seed_library_item(self, client, label):
|
|
seed = self.make_temp(self.uploader, label)
|
|
response = self.resolve_bulk(client, [seed.id], "s")
|
|
self.assertEqual(response.status_code, 200)
|
|
return seed
|
|
|
|
def test_upload_defers_visual_similarity_to_its_phase(self):
|
|
client = self.api_client(self.uploader)
|
|
self.seed_library_item(client, "seed-defer")
|
|
response = self.upload_via_api(client)
|
|
self.assertEqual(response.status_code, 201)
|
|
body = response.json()
|
|
# Uploading must not run the expensive library scan.
|
|
self.assertEqual(body["status"], TempUpload.STATUS_PENDING)
|
|
self.assertFalse(body["visual_matches"])
|
|
|
|
def test_visual_match_phase_flags_similar_library_items(self):
|
|
client = self.api_client(self.uploader)
|
|
self.seed_library_item(client, "seed-visual")
|
|
temp = self.make_temp(self.uploader, "check-visual")
|
|
response = client.post(f"/api/uploads/{temp.id}/visual-match/")
|
|
self.assertEqual(response.status_code, 200)
|
|
body = response.json()
|
|
self.assertEqual(body["status"], TempUpload.STATUS_VISUAL_MATCH)
|
|
self.assertGreaterEqual(len(body["visual_matches"]), 1)
|
|
temp.refresh_from_db()
|
|
self.assertEqual(temp.status, TempUpload.STATUS_VISUAL_MATCH)
|
|
|
|
def test_visual_match_phase_with_empty_library_stays_pending(self):
|
|
client = self.api_client(self.uploader)
|
|
temp = self.make_temp(self.uploader, "no-matches")
|
|
response = client.post(f"/api/uploads/{temp.id}/visual-match/")
|
|
self.assertEqual(response.status_code, 200)
|
|
body = response.json()
|
|
self.assertEqual(body["visual_matches"], [])
|
|
self.assertEqual(body["status"], TempUpload.STATUS_PENDING)
|
|
|
|
def test_visual_match_phase_rejects_completed_uploads(self):
|
|
client = self.api_client(self.uploader)
|
|
temp = self.make_temp(
|
|
self.uploader, "done-visual", status=TempUpload.STATUS_COMPLETED
|
|
)
|
|
response = client.post(f"/api/uploads/{temp.id}/visual-match/")
|
|
self.assertEqual(response.status_code, 400)
|
|
|
|
|
|
class IqdbRecordingTests(TestCase):
|
|
"""The modal needs to tell "checked, no match" from "never checked"."""
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
super().setUpClass()
|
|
cls._tmp = tempfile.mkdtemp(prefix="j621-iqdb-")
|
|
cls._settings = override_settings(MEDIA_ROOT=cls._tmp)
|
|
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="iqdb-uploader", password="iqdb-pass-123456"
|
|
)
|
|
self.uploader.role = "uploader"
|
|
self.uploader.save(update_fields=["role"])
|
|
self.client = Client()
|
|
self.client.defaults["HTTP_AUTHORIZATION"] = (
|
|
f"Token {Token.objects.create(user=self.uploader).key}"
|
|
)
|
|
|
|
def make_temp(self):
|
|
return TempUpload.objects.create(
|
|
user=self.uploader,
|
|
file=SimpleUploadedFile(
|
|
"checked.png", TINY_PNG, content_type="image/png"
|
|
),
|
|
original_filename="checked.png",
|
|
md5=hashlib.md5(b"checked").hexdigest(),
|
|
size=len(TINY_PNG),
|
|
)
|
|
|
|
def test_empty_result_records_the_check_without_a_match(self):
|
|
temp = self.make_temp()
|
|
response = jpost(
|
|
self.client, f"/api/uploads/{temp.id}/iqdb/", {"results": []}
|
|
)
|
|
self.assertEqual(response.status_code, 200)
|
|
temp.refresh_from_db()
|
|
self.assertEqual(temp.iqdb_data, [])
|
|
# No candidates must not masquerade as a visual match.
|
|
self.assertEqual(temp.status, TempUpload.STATUS_PENDING)
|
|
|
|
def test_candidates_are_stored_and_flag_a_visual_match(self):
|
|
temp = self.make_temp()
|
|
response = jpost(
|
|
self.client,
|
|
f"/api/uploads/{temp.id}/iqdb/",
|
|
{
|
|
"results": [
|
|
{
|
|
"post_id": 123,
|
|
"score": 91.5,
|
|
"preview_url": "https://static1.e621.net/data/preview/ab/cd/x.jpg",
|
|
"rating": "q",
|
|
"md5": "a" * 32,
|
|
"score_total": 12,
|
|
"fav_count": 3,
|
|
"width": 800,
|
|
"height": 600,
|
|
"tags_preview": ["canine", "solo"],
|
|
}
|
|
]
|
|
},
|
|
)
|
|
self.assertEqual(response.status_code, 200)
|
|
temp.refresh_from_db()
|
|
self.assertEqual(len(temp.iqdb_data), 1)
|
|
self.assertEqual(temp.iqdb_data[0]["post_id"], 123)
|
|
self.assertEqual(temp.status, TempUpload.STATUS_VISUAL_MATCH)
|
|
|
|
def test_rejects_a_non_list_payload(self):
|
|
temp = self.make_temp()
|
|
response = jpost(
|
|
self.client, f"/api/uploads/{temp.id}/iqdb/", {"results": "nope"}
|
|
)
|
|
self.assertEqual(response.status_code, 400)
|