Files
vision/tests/test_image_io.py

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