diff --git a/tests/unit/test_pagination.py b/tests/unit/test_pagination.py new file mode 100755 index 000000000..37ee01da8 --- /dev/null +++ b/tests/unit/test_pagination.py @@ -0,0 +1,337 @@ +""" +pagination 通用分页器单元测试 + +覆盖: +- PaginationParams: 默认值/边界/校验/offset/limit +- PaginationMeta: from_params 各种边界场景 +- PaginatedResponse: create 工厂方法 +- paginate: 内存分页函数 +""" + +import pytest +from pydantic import ValidationError + +from packages.application.common.pagination import ( + PaginatedResponse, + PaginationMeta, + PaginationParams, + paginate, +) + +# ============================================================ +# PaginationParams +# ============================================================ + + +class TestPaginationParamsDefaults: + """默认值测试""" + + def test_default_page_is_1(self): + params = PaginationParams() + assert params.page == 1 + + def test_default_page_size_is_20(self): + params = PaginationParams() + assert params.page_size == 20 + + def test_default_offset_is_0(self): + params = PaginationParams() + assert params.offset == 0 + + def test_default_limit_is_20(self): + params = PaginationParams() + assert params.limit == 20 + + +class TestPaginationParamsValidation: + """参数校验""" + + @pytest.mark.parametrize("page", [1, 2, 100, 9999]) + def test_valid_page_values(self, page): + params = PaginationParams(page=page) + assert params.page == page + + def test_page_zero_raises(self): + with pytest.raises(ValidationError): + PaginationParams(page=0) + + def test_page_negative_raises(self): + with pytest.raises(ValidationError): + PaginationParams(page=-1) + + @pytest.mark.parametrize("page_size", [1, 20, 50, 100]) + def test_valid_page_size_values(self, page_size): + params = PaginationParams(page_size=page_size) + assert params.page_size == page_size + + def test_page_size_zero_raises(self): + with pytest.raises(ValidationError): + PaginationParams(page_size=0) + + def test_page_size_negative_raises(self): + with pytest.raises(ValidationError): + PaginationParams(page_size=-5) + + def test_page_size_over_100_raises(self): + with pytest.raises(ValidationError): + PaginationParams(page_size=101) + + def test_invalid_page_type_raises(self): + with pytest.raises(ValidationError): + PaginationParams(page="abc") + + def test_invalid_page_size_type_raises(self): + with pytest.raises(ValidationError): + PaginationParams(page_size="abc") + + +class TestPaginationParamsOffset: + """offset 属性计算""" + + def test_page_1_offset_0(self): + params = PaginationParams(page=1, page_size=20) + assert params.offset == 0 + + def test_page_2_offset_page_size(self): + params = PaginationParams(page=2, page_size=20) + assert params.offset == 20 + + def test_page_3_offset_2x_page_size(self): + params = PaginationParams(page=3, page_size=20) + assert params.offset == 40 + + def test_page_5_page_size_10_offset_40(self): + params = PaginationParams(page=5, page_size=10) + assert params.offset == 40 + + def test_page_1_page_size_100_offset_0(self): + params = PaginationParams(page=1, page_size=100) + assert params.offset == 0 + + +class TestPaginationParamsLimit: + """limit 属性""" + + def test_limit_equals_page_size(self): + params = PaginationParams(page_size=20) + assert params.limit == 20 + + def test_limit_1(self): + params = PaginationParams(page_size=1) + assert params.limit == 1 + + def test_limit_100(self): + params = PaginationParams(page_size=100) + assert params.limit == 100 + + +# ============================================================ +# PaginationMeta.from_params +# ============================================================ + + +class TestPaginationMetaFromParams: + """from_params 工厂方法""" + + def test_empty_total_zero(self): + params = PaginationParams(page=1, page_size=20) + meta = PaginationMeta.from_params(params, total=0) + assert meta.total == 0 + assert meta.total_pages == 0 + assert meta.has_next is False + assert meta.has_prev is False + + def test_exactly_one_page(self): + params = PaginationParams(page=1, page_size=20) + meta = PaginationMeta.from_params(params, total=20) + assert meta.total_pages == 1 + assert meta.has_next is False + assert meta.has_prev is False + + def test_less_than_one_page(self): + params = PaginationParams(page=1, page_size=20) + meta = PaginationMeta.from_params(params, total=15) + assert meta.total_pages == 1 + assert meta.has_next is False + assert meta.has_prev is False + + def test_multiple_pages_first_page(self): + params = PaginationParams(page=1, page_size=20) + meta = PaginationMeta.from_params(params, total=50) + assert meta.total_pages == 3 + assert meta.has_next is True + assert meta.has_prev is False + + def test_multiple_pages_middle_page(self): + params = PaginationParams(page=2, page_size=20) + meta = PaginationMeta.from_params(params, total=50) + assert meta.total_pages == 3 + assert meta.has_next is True + assert meta.has_prev is True + + def test_multiple_pages_last_page(self): + params = PaginationParams(page=3, page_size=20) + meta = PaginationMeta.from_params(params, total=50) + assert meta.total_pages == 3 + assert meta.has_next is False + assert meta.has_prev is True + + def test_exact_division(self): + params = PaginationParams(page=2, page_size=20) + meta = PaginationMeta.from_params(params, total=40) + assert meta.total_pages == 2 + assert meta.has_next is False + assert meta.has_prev is True + + def test_non_exact_division_ceil(self): + params = PaginationParams(page=1, page_size=20) + meta = PaginationMeta.from_params(params, total=41) + assert meta.total_pages == 3 + + def test_total_1_page_size_20(self): + params = PaginationParams(page=1, page_size=20) + meta = PaginationMeta.from_params(params, total=1) + assert meta.total_pages == 1 + assert meta.has_next is False + assert meta.has_prev is False + + def test_page_beyond_total_pages(self): + params = PaginationParams(page=10, page_size=20) + meta = PaginationMeta.from_params(params, total=50) + assert meta.total_pages == 3 + assert meta.has_next is False + assert meta.has_prev is True + + def test_preserves_params_values(self): + params = PaginationParams(page=3, page_size=15) + meta = PaginationMeta.from_params(params, total=100) + assert meta.page == 3 + assert meta.page_size == 15 + assert meta.total == 100 + + +# ============================================================ +# PaginatedResponse.create +# ============================================================ + + +class TestPaginatedResponseCreate: + """create 工厂方法""" + + def test_create_with_data(self): + params = PaginationParams(page=1, page_size=20) + data = [1, 2, 3] + response = PaginatedResponse.create(data, params, total=100) + assert response.data == data + assert response.pagination.total == 100 + assert response.pagination.page == 1 + assert response.pagination.page_size == 20 + + def test_create_with_empty_data(self): + params = PaginationParams(page=1, page_size=20) + response = PaginatedResponse.create([], params, total=0) + assert response.data == [] + assert response.pagination.total == 0 + assert response.pagination.total_pages == 0 + + def test_create_preserves_list_type(self): + params = PaginationParams(page=1, page_size=20) + data = ["a", "b", "c"] + response = PaginatedResponse.create(data, params, total=10) + assert response.data == ["a", "b", "c"] + assert len(response.data) == 3 + + +# ============================================================ +# paginate 函数 +# ============================================================ + + +class TestPaginateFunction: + """内存分页函数""" + + def test_empty_list(self): + params = PaginationParams(page=1, page_size=20) + result = paginate([], params) + assert result.data == [] + assert result.pagination.total == 0 + assert result.pagination.total_pages == 0 + + def test_first_page(self): + items = list(range(50)) + params = PaginationParams(page=1, page_size=20) + result = paginate(items, params) + assert result.data == list(range(20)) + assert result.pagination.total == 50 + assert result.pagination.total_pages == 3 + assert result.pagination.has_next is True + assert result.pagination.has_prev is False + + def test_middle_page(self): + items = list(range(50)) + params = PaginationParams(page=2, page_size=20) + result = paginate(items, params) + assert result.data == list(range(20, 40)) + assert result.pagination.has_next is True + assert result.pagination.has_prev is True + + def test_last_page(self): + items = list(range(50)) + params = PaginationParams(page=3, page_size=20) + result = paginate(items, params) + assert result.data == list(range(40, 50)) + assert len(result.data) == 10 + assert result.pagination.has_next is False + assert result.pagination.has_prev is True + + def test_page_beyond_total(self): + items = list(range(25)) + params = PaginationParams(page=10, page_size=20) + result = paginate(items, params) + assert result.data == [] + assert result.pagination.total == 25 + assert result.pagination.total_pages == 2 + + def test_page_size_larger_than_total(self): + items = list(range(5)) + params = PaginationParams(page=1, page_size=20) + result = paginate(items, params) + assert result.data == items + assert result.pagination.total_pages == 1 + assert result.pagination.has_next is False + + def test_single_item(self): + items = [42] + params = PaginationParams(page=1, page_size=20) + result = paginate(items, params) + assert result.data == [42] + assert result.pagination.total == 1 + + def test_page_size_1(self): + items = list(range(5)) + params = PaginationParams(page=3, page_size=1) + result = paginate(items, params) + assert result.data == [2] + assert result.pagination.total_pages == 5 + + def test_exact_page_size(self): + items = list(range(40)) + params = PaginationParams(page=2, page_size=20) + result = paginate(items, params) + assert result.data == list(range(20, 40)) + assert result.pagination.total_pages == 2 + assert result.pagination.has_next is False + + def test_string_items(self): + items = ["a", "b", "c", "d", "e"] + params = PaginationParams(page=2, page_size=2) + result = paginate(items, params) + assert result.data == ["c", "d"] + assert result.pagination.total == 5 + + def test_does_not_mutate_original_list(self): + items = list(range(10)) + original = items.copy() + params = PaginationParams(page=1, page_size=3) + paginate(items, params) + assert items == original diff --git a/tests/unit/test_url_security.py b/tests/unit/test_url_security.py index 63da47669..20bdcb026 100755 --- a/tests/unit/test_url_security.py +++ b/tests/unit/test_url_security.py @@ -1,296 +1,597 @@ -"""URL 安全校验工具单元测试 — SSRF 防护.""" +""" +url_security URL安全校验单元测试 -from __future__ import annotations +覆盖: +- validate_url_safety: scheme/主机/端口/SSRF/内网域名/白名单 +- is_url_safe: 便捷函数 +- UrlSecurityError / NoRedirectHandler +- _validate_magic_number: 文件魔数校验 +- safe_download_file / safe_download_bytes: mock 网络测试 +""" import os -import sys -import unittest - -sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker")) - -import shutil import tempfile +from unittest.mock import MagicMock, patch -from video_processing.url_security import ( # noqa: E402 +import pytest + +from packages.shared.url_security import ( ALLOWED_AUDIO_MIME_TYPES, + ALLOWED_IMAGE_MIME_TYPES, + ALLOWED_PORTS, + ALLOWED_SCHEMES, + MAX_URL_LENGTH, + NoRedirectHandler, UrlSecurityError, + _check_internal_hostnames, + _is_trusted_domain, + _validate_magic_number, is_url_safe, safe_download_bytes, safe_download_file, validate_url_safety, ) - -class TestUrlSecurityValidation(unittest.TestCase): - """URL 安全校验测试.""" - - # ── Scheme 白名单 ────────────────────────────────────────────────────── - - def test_http_scheme_allowed(self): - """HTTP scheme 应该被允许.""" - result = validate_url_safety("http://example.com/test", purpose="test") - self.assertEqual(result, "http://example.com/test") - - def test_https_scheme_allowed(self): - """HTTPS scheme 应该被允许.""" - result = validate_url_safety("https://example.com/test", purpose="test") - self.assertEqual(result, "https://example.com/test") - - def test_file_scheme_rejected(self): - """file:// scheme 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("file:///etc/passwd", purpose="test") - - def test_ftp_scheme_rejected(self): - """ftp:// scheme 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("ftp://example.com/test", purpose="test") - - def test_empty_scheme_rejected(self): - """空 scheme 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("example.com/test", purpose="test") - - # ── 端口白名单 ──────────────────────────────────────────────────────── - - def test_port_80_allowed(self): - """端口 80 应该被允许.""" - # 80端口是默认HTTP端口,不显式指定也可以 - result = validate_url_safety("http://example.com:80/test", purpose="test") - self.assertIn("example.com", result) - - def test_port_443_allowed(self): - """端口 443 应该被允许.""" - result = validate_url_safety("https://example.com:443/test", purpose="test") - self.assertIn("example.com", result) - - def test_port_8080_rejected(self): - """非标准端口 8080 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://example.com:8080/test", purpose="test") - - def test_port_22_rejected(self): - """SSH 端口 22 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://example.com:22/test", purpose="test") - - # ── SSRF: 直接 IP 访问 ─────────────────────────────────────────────── - - def test_loopback_ip_rejected(self): - """回环地址 127.0.0.1 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://127.0.0.1/test", purpose="test") - - def test_private_ip_192_rejected(self): - """内网地址 192.168.x.x 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://192.168.1.1/test", purpose="test") - - def test_private_ip_10_rejected(self): - """内网地址 10.x.x.x 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://10.0.0.1/test", purpose="test") - - def test_private_ip_172_rejected(self): - """内网地址 172.16.x.x 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://172.16.0.1/test", purpose="test") - - def test_unspecified_ip_rejected(self): - """未指定地址 0.0.0.0 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://0.0.0.0/test", purpose="test") - - def test_ipv6_loopback_rejected(self): - """IPv6 回环 ::1 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://[::1]/test", purpose="test") - - def test_ipv6_link_local_rejected(self): - """IPv6 链路本地地址应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://[fe80::1]/test", purpose="test") - - # ── SSRF: 内网主机名 ───────────────────────────────────────────────── - - def test_localhost_rejected(self): - """localhost 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://localhost/test", purpose="test") - - def test_local_domain_rejected(self): - """.local 域名应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://printer.local/test", purpose="test") - - def test_internal_domain_rejected(self): - """.internal 域名应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http://db.internal/test", purpose="test") - - # ── URL 格式校验 ───────────────────────────────────────────────────── - - def test_empty_url_rejected(self): - """空 URL 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("", purpose="test") - - def test_none_url_rejected(self): - """None URL 应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety(None, purpose="test") # type: ignore - - def test_url_too_long_rejected(self): - """超长 URL 应该被拒绝.""" - long_url = "https://example.com/" + "a" * 3000 - with self.assertRaises(UrlSecurityError): - validate_url_safety(long_url, purpose="test") - - def test_no_hostname_rejected(self): - """缺少主机名应该被拒绝.""" - with self.assertRaises(UrlSecurityError): - validate_url_safety("http:///test", purpose="test") - - # ── is_url_safe 便捷函数 ───────────────────────────────────────────── - - def test_is_url_safe_true(self): - """安全 URL 应该返回 True.""" - self.assertTrue(is_url_safe("https://example.com/test", purpose="test")) - - def test_is_url_safe_false(self): - """不安全 URL 应该返回 False.""" - self.assertFalse(is_url_safe("http://127.0.0.1/test", purpose="test")) - - def test_is_url_safe_empty(self): - """空 URL 应该返回 False.""" - self.assertFalse(is_url_safe("", purpose="test")) +# ── validate_url_safety 基础校验 ───────────────────────────────────────────── -if __name__ == "__main__": - unittest.main() +class TestValidateUrlSafetyBasics: + """URL 安全校验基础测试""" + + def test_valid_http_url(self): + url = "http://example.com/file.mp4" + result = validate_url_safety(url) + assert result == url + + def test_valid_https_url(self): + url = "https://example.com/file.mp4" + result = validate_url_safety(url) + assert result == url + + def test_empty_url_raises(self): + with pytest.raises(UrlSecurityError, match="为空"): + validate_url_safety("") + + def test_none_url_raises(self): + with pytest.raises(UrlSecurityError): + validate_url_safety(None) + + def test_url_too_long_raises(self): + long_url = "https://example.com/" + "a" * 2050 + with pytest.raises(UrlSecurityError, match="过长"): + validate_url_safety(long_url) + + def test_url_at_max_length_ok(self): + base = "https://example.com/" + pad = "a" * (MAX_URL_LENGTH - len(base)) + url = base + pad + assert len(url) <= MAX_URL_LENGTH + result = validate_url_safety(url) + assert result == url + + def test_invalid_scheme_ftp_raises(self): + with pytest.raises(UrlSecurityError, match="scheme"): + validate_url_safety("ftp://example.com/file") + + def test_invalid_scheme_file_raises(self): + with pytest.raises(UrlSecurityError, match="scheme"): + validate_url_safety("file:///etc/passwd") + + def test_invalid_scheme_data_raises(self): + with pytest.raises(UrlSecurityError, match="scheme"): + validate_url_safety("data:text/html,