"""Staged uploads: complete-set listing, bulk rating and IQDB recording.""" import base64 import hashlib 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 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 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 BulkResolveTests(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) 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)