from __future__ import annotations import socket import unittest from unittest.mock import Mock, patch import requests from common.image_io import ImageLoadError, bytes_to_pil, fetch_url_bytes class _Response: def __init__(self, status_code=200, *, headers=None, body=b"image-bytes"): self.status_code = status_code self.headers = headers or {"content-type": "image/webp"} self.body = body def __enter__(self): return self def __exit__(self, *_args): return False def iter_content(self, chunk_size=0): yield self.body def raise_for_status(self): if self.status_code >= 400: raise requests.HTTPError(f"HTTP {self.status_code}") class ImageIoClassificationTests(unittest.TestCase): def assert_category(self, operation, category): with self.assertRaises(ImageLoadError) as raised: operation() self.assertEqual(raised.exception.category, category) @patch("common.image_io.socket.getaddrinfo") def test_invalid_url_scheme(self, _getaddrinfo): self.assert_category(lambda: fetch_url_bytes("ftp://cdn.example/a.webp"), "invalid_scheme") @patch("common.image_io.socket.getaddrinfo") def test_missing_hostname(self, _getaddrinfo): self.assert_category(lambda: fetch_url_bytes("https:///a.webp"), "missing_hostname") @patch("common.image_io.socket.getaddrinfo", side_effect=socket.gaierror("NXDOMAIN")) def test_dns_failure(self, _getaddrinfo): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "dns_failure") @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("127.0.0.1", 443))]) def test_private_address_rejected(self, _getaddrinfo): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "blocked_address") @patch("common.image_io.requests.get", side_effect=requests.ConnectTimeout("connect")) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_connect_timeout(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "connect_timeout") @patch("common.image_io.requests.get", side_effect=requests.ReadTimeout("read")) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_read_timeout(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "read_timeout") @patch("common.image_io.requests.get", side_effect=requests.ConnectionError("connection")) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_generic_connection_failure(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "connection_error") @patch("common.image_io.requests.get", return_value=_Response(302)) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_redirect_without_location(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "redirect_missing_location") @patch("common.image_io.requests.get", side_effect=[_Response(302, headers={"location": "https://cdn.example/a.webp"})] * 4) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_too_many_redirects(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "too_many_redirects") @patch("common.image_io.requests.get", return_value=_Response(404)) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_remote_404(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "remote_http_4xx") @patch("common.image_io.requests.get", return_value=_Response(500)) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_remote_500(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "remote_http_5xx") @patch("common.image_io.requests.get", return_value=_Response(headers={"content-type": "text/html"})) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_invalid_content_type(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp"), "invalid_content_type") @patch("common.image_io.requests.get", return_value=_Response(body=b"12345")) @patch("common.image_io.socket.getaddrinfo", return_value=[(None, None, None, None, ("93.184.216.34", 443))]) def test_max_bytes_exceeded(self, _getaddrinfo, _get): self.assert_category(lambda: fetch_url_bytes("https://cdn.example/a.webp", max_bytes=4), "image_too_large") def test_decode_failure(self): self.assert_category(lambda: bytes_to_pil(b"not-an-image"), "decode_failure") if __name__ == "__main__": unittest.main()