Files
J621/backend/apps/library/tests/test_uploads.py
T
JakeBreath 3a07481dfc Wait for all uploads, then batch MD5 -> visual -> IQDB with bulk links
Batching must not start while files are still being uploaded, and the MD5
phase must move a whole chunk at once instead of one resolve per file:

- the upload queue drains completely first (failed uploads included) before
  any matching starts;
- phase 1 asks e621 for every md5 (75 per posts.json request), builds the
  md5 -> post map from the response, and sends the matches to the new
  POST /api/uploads/link-bulk/ action, so a whole 75-file chunk moves into
  Indexed in a single board update;
- link-bulk indexes the staged file directly when the post's MD5 matches
  (identical bytes), so there is no per-file download round trip;
- phase 2 runs local visual similarity for whatever stayed pending, phase 3
  the IQDB queue.

Verified end to end with real e621 files: one md5 query for the batch, one
link-bulk call, both matching files flipped to Indexed together, then the
visual and IQDB phases. 23 library tests green (link-bulk, visual phase,
deferred visual matching).
2026-09-19 11:10:24 -05:00

403 lines
15 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, payload=None
):
return TempUpload.objects.create(
user=user,
file=SimpleUploadedFile(
f"{label}.png", payload or TINY_PNG, content_type="image/png"
),
original_filename=f"{label}.png",
md5=hashlib.md5(label.encode()).hexdigest(),
size=len(payload or 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)
def test_link_bulk_moves_a_batch_with_one_call(self):
client = self.api_client(self.uploader)
first = self.make_temp(self.uploader, "bulk-link-1")
second = self.make_temp(
self.uploader, "bulk-link-2", payload=similar_png((10, 200, 30))
)
response = jpost(
client,
"/api/uploads/link-bulk/",
{
"links": [
{
"temp_id": str(first.id),
"post": {
"id": 900001,
"rating": "s",
"file": {
"md5": first.md5,
# Identical MD5: the staged file is indexed
# without downloading the URL.
"url": "https://static1.e621.net/data/fake-1.png",
},
},
},
{
"temp_id": str(second.id),
"post": {
"id": 900002,
"rating": "q",
"file": {
"md5": second.md5,
"url": "https://static1.e621.net/data/fake-2.png",
},
},
},
]
},
)
self.assertEqual(response.status_code, 200)
body = response.json()
self.assertEqual(body["errors"], [])
self.assertEqual(len(body["updated"]), 2)
for temp, post_id in ((first, 900001), (second, 900002)):
temp.refresh_from_db()
self.assertEqual(temp.status, TempUpload.STATUS_COMPLETED)
self.assertEqual(temp.library_item_id is not None, True)
self.assertEqual(temp.e621_post_id, post_id)
self.assertEqual(temp.resolution, TempUpload.RESOLUTION_AUTO_MD5)
self.assertEqual(temp.library_item.e621_post_id, post_id)
self.assertEqual(MediaItem.objects.count(), 2)
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)