Files
J621/backend/apps/library/tests/test_uploads.py
T
2026-09-21 09:01:01 -05:00

698 lines
26 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 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 import upload_pipeline
from apps.library.models import MediaItem, TempUpload, UploadRun
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)
@override_settings(UPLOAD_PIPELINE_AUTOSTART=False)
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)
@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)