Files
2026-09-21 09:01:01 -05:00

100 lines
3.6 KiB
Python

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