Upload updates
This commit is contained in:
@@ -0,0 +1,99 @@
|
||||
"""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")
|
||||
Reference in New Issue
Block a user