100 lines
3.6 KiB
Python
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")
|