"""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)