"""The server-side e621 client: batching, retries and IQDB stream handling.""" from pathlib import Path from tempfile import TemporaryDirectory from unittest import mock from django.test import SimpleTestCase from apps.library import e621 class FakeResponse: def __init__(self, status_code, payload=None, headers=None): self.status_code = status_code self._payload = payload self.headers = headers or {} def json(self): if self._payload is None: raise ValueError("not json") return self._payload class E621ClientTests(SimpleTestCase): def test_iqdb_search_reopens_the_file_on_retry(self): bodies = [] def fake_request(method, url, **kwargs): handle = kwargs["files"]["search[file]"][1] bodies.append(handle.read()) if len(bodies) == 1: return FakeResponse( 429, {"success": False, "message": "Throttled"} ) return FakeResponse(200, [{"post_id": 1, "score": 90.0}]) with TemporaryDirectory() as tmp: path = Path(tmp) / "x.png" path.write_bytes(b"image-bytes") with mock.patch.object( e621.requests, "request", side_effect=fake_request ), mock.patch.object(e621.time, "sleep"), mock.patch.object( e621, "_wait_for_slot" ): result = e621.iqdb_search(None, path) # Both attempts must send the full body, not the consumed handle. self.assertEqual(bodies, [b"image-bytes", b"image-bytes"]) self.assertEqual(result, [{"post_id": 1, "score": 90.0}]) def test_check_md5_batch_keys_by_md5(self): payload = { "posts": [ {"id": 5, "file": {"md5": "a" * 32}}, {"id": 6, "file": {"md5": "b" * 32}}, ] } with mock.patch.object(e621, "get", return_value=payload) as getter: found = e621.check_md5_batch(None, ["A" * 32, "b" * 32]) self.assertEqual(set(found), {"a" * 32, "b" * 32}) params = getter.call_args.kwargs["params"] self.assertTrue(params["tags"].startswith("md5:")) self.assertEqual(params["limit"], 2) def test_auth_errors_are_not_retried(self): calls = [] def fake_request(*args, **kwargs): calls.append(1) return FakeResponse(403, {"error": "nope"}) with mock.patch.object( e621.requests, "request", side_effect=fake_request ), mock.patch.object(e621, "_wait_for_slot"): with self.assertRaises(e621.E621AuthError): e621._request(None, "GET", "/posts.json", require_auth=False) self.assertEqual(len(calls), 1) def test_load_shedding_html_raises_rate_limited_after_retries(self): with mock.patch.object( e621.requests, "request", return_value=FakeResponse(200, None), ), mock.patch.object(e621.time, "sleep"), mock.patch.object( e621, "_wait_for_slot" ): with self.assertRaises(e621.E621RateLimited): e621._request( None, "GET", "/posts.json", require_auth=False, attempts=2 ) def test_fetch_posts_by_ids_queries_with_id_tag(self): with mock.patch.object( e621, "get", return_value={"posts": [{"id": 9}]} ) as getter: posts = e621.fetch_posts_by_ids(None, [9]) self.assertEqual(posts, [{"id": 9}]) params = getter.call_args.kwargs["params"] self.assertEqual(params["tags"], "id:9")