105 lines
5.2 KiB
Python
105 lines
5.2 KiB
Python
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()
|