diff --git a/tests/unit/domain/test_url_security.py b/tests/unit/domain/test_url_security.py new file mode 100755 index 000000000..6a22c222a --- /dev/null +++ b/tests/unit/domain/test_url_security.py @@ -0,0 +1,567 @@ +"""url_security 单测. + +domain 层 URL 安全校验纯逻辑模块,0 网络依赖。 +覆盖 SSRF 防护、主机名校验、IP 检查、魔数校验等。 +""" + +from __future__ import annotations + +import pytest + +from packages.domain.url_security import ( + ALLOWED_IMAGE_MIME_TYPES, + ALLOWED_PORTS, + ALLOWED_SCHEMES, + ALLOWED_VIDEO_MIME_TYPES, + MAGIC_NUMBERS, + MAX_URL_LENGTH, + UrlSecurityError, + check_internal_hostname, + check_ssrf_ip, + is_ip_address, + is_trusted_domain, + is_url_basic_safe, + validate_magic_number, + validate_url_basic, +) + +# ═══════════════════════════════════════════════════════════════════════════════ +# 常量与异常类 +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestConstants: + """常量测试.""" + + def test_allowed_schemes(self): + """允许的 scheme 包含 http 和 https.""" + assert "http" in ALLOWED_SCHEMES + assert "https" in ALLOWED_SCHEMES + + def test_allowed_ports(self): + """允许的端口:80, 443.""" + assert 80 in ALLOWED_PORTS + assert 443 in ALLOWED_PORTS + + def test_max_url_length(self): + """最大 URL 长度 2048.""" + assert MAX_URL_LENGTH == 2048 + + def test_magic_numbers_has_common_formats(self): + """魔数表包含常见格式.""" + assert "image/jpeg" in MAGIC_NUMBERS + assert "image/png" in MAGIC_NUMBERS + assert "image/gif" in MAGIC_NUMBERS + assert "video/mp4" in MAGIC_NUMBERS + assert "audio/mpeg" in MAGIC_NUMBERS + + +class TestUrlSecurityError: + """异常类测试.""" + + def test_is_value_error(self): + """UrlSecurityError 继承 ValueError.""" + assert issubclass(UrlSecurityError, ValueError) + + def test_raise_with_message(self): + """抛出时携带错误信息.""" + with pytest.raises(UrlSecurityError, match="test error"): + raise UrlSecurityError("test error") + + +# ═══════════════════════════════════════════════════════════════════════════════ +# check_internal_hostname +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestCheckInternalHostname: + """内部主机名检查测试.""" + + def test_normal_domain_passes(self): + """普通外部域名通过.""" + check_internal_hostname("example.com") + check_internal_hostname("www.google.com") + + def test_localhost_blocked(self): + """localhost 被拦截.""" + with pytest.raises(UrlSecurityError, match="内部主机名"): + check_internal_hostname("localhost") + + def test_localhost_case_insensitive(self): + """大小写不敏感.""" + with pytest.raises(UrlSecurityError): + check_internal_hostname("LOCALHOST") + with pytest.raises(UrlSecurityError): + check_internal_hostname("LocalHost") + + def test_localhost_localdomain_blocked(self): + """localhost.localdomain 被拦截.""" + with pytest.raises(UrlSecurityError): + check_internal_hostname("localhost.localdomain") + + def test_metadata_blocked(self): + """metadata 被拦截.""" + with pytest.raises(UrlSecurityError): + check_internal_hostname("metadata") + + def test_metadata_google_internal_blocked(self): + """GCP 元数据服务被拦截.""" + with pytest.raises(UrlSecurityError): + check_internal_hostname("metadata.google.internal") + + def test_cloud_metadata_ip_blocked(self): + """云元数据 IP 169.254.169.254 被拦截.""" + with pytest.raises(UrlSecurityError): + check_internal_hostname("169.254.169.254") + + def test_local_suffix_blocked(self): + """.local 后缀域名被拦截.""" + with pytest.raises(UrlSecurityError, match="内网域名"): + check_internal_hostname("myhost.local") + + def test_internal_suffix_blocked(self): + """.internal 后缀被拦截.""" + with pytest.raises(UrlSecurityError): + check_internal_hostname("svc.cluster.internal") + + def test_localdomain_suffix_blocked(self): + """.localdomain 后缀被拦截.""" + with pytest.raises(UrlSecurityError): + check_internal_hostname("host.localdomain") + + def test_com_domain_not_blocked(self): + """.com 域名不被拦截.""" + check_internal_hostname("example.com") + check_internal_hostname("sub.example.com") + + def test_subdomain_of_public_domain_ok(self): + """公网域名的子域名正常.""" + check_internal_hostname("api.example.com") + check_internal_hostname("cdn.assets.example.org") + + +# ═══════════════════════════════════════════════════════════════════════════════ +# is_trusted_domain +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestIsTrustedDomain: + """可信域名匹配测试.""" + + def test_empty_trusted_domains_allows_all(self): + """空集合允许所有域名.""" + assert is_trusted_domain("anything.com", set()) is True + assert is_trusted_domain("anywhere.org", set()) is True + + def test_exact_match(self): + """精确匹配.""" + trusted = {"example.com", "example.org"} + assert is_trusted_domain("example.com", trusted) is True + assert is_trusted_domain("example.org", trusted) is True + + def test_subdomain_match(self): + """子域名匹配.""" + trusted = {"example.com"} + assert is_trusted_domain("api.example.com", trusted) is True + assert is_trusted_domain("cdn.assets.example.com", trusted) is True + + def test_no_match(self): + """不匹配.""" + trusted = {"example.com"} + assert is_trusted_domain("other.com", trusted) is False + assert is_trusted_domain("example.net", trusted) is False + + def test_case_insensitive(self): + """大小写不敏感.""" + trusted = {"Example.COM"} + assert is_trusted_domain("example.com", trusted) is True + assert is_trusted_domain("API.EXAMPLE.COM", trusted) is True + + def test_partial_match_no(self): + """域名部分相同但不是子域名不匹配.""" + trusted = {"example.com"} + # fakeexample.com 不是 example.com 的子域名 + assert is_trusted_domain("fakeexample.com", trusted) is False + + def test_none_trusted_domains(self): + """trusted_domains 为 None 时由调用方处理,空 set 全允许.""" + # 传空集合时全允许 + assert is_trusted_domain("a.com", set()) is True + + +# ═══════════════════════════════════════════════════════════════════════════════ +# check_ssrf_ip +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestCheckSsrIp: + """IP SSRF 检查测试.""" + + def test_public_ip_passes(self): + """公网 IP 通过.""" + check_ssrf_ip("8.8.8.8") + check_ssrf_ip("1.1.1.1") + check_ssrf_ip("114.114.114.114") + + def test_loopback_blocked(self): + """回环地址被拦截.""" + with pytest.raises(UrlSecurityError, match="回环"): + check_ssrf_ip("127.0.0.1") + with pytest.raises(UrlSecurityError): + check_ssrf_ip("127.0.0.53") + + def test_private_ip_blocked(self): + """私有内网 IP 被拦截.""" + with pytest.raises(UrlSecurityError, match="内网"): + check_ssrf_ip("192.168.1.1") + with pytest.raises(UrlSecurityError): + check_ssrf_ip("10.0.0.1") + with pytest.raises(UrlSecurityError): + check_ssrf_ip("172.16.0.1") + + def test_link_local_blocked(self): + """链路本地地址被拦截.""" + with pytest.raises(UrlSecurityError, match="链路本地"): + check_ssrf_ip("169.254.169.254") + with pytest.raises(UrlSecurityError): + check_ssrf_ip("169.254.1.1") + + def test_multicast_blocked(self): + """组播地址被拦截.""" + with pytest.raises(UrlSecurityError, match="组播"): + check_ssrf_ip("224.0.0.1") + with pytest.raises(UrlSecurityError): + check_ssrf_ip("239.255.255.250") + + def test_unspecified_blocked(self): + """未指定地址被拦截.""" + with pytest.raises(UrlSecurityError, match="未指定"): + check_ssrf_ip("0.0.0.0") + + def test_ipv6_loopback_blocked(self): + """IPv6 回环地址被拦截.""" + with pytest.raises(UrlSecurityError): + check_ssrf_ip("::1") + + def test_ipv6_private_blocked(self): + """IPv6 内网地址被拦截.""" + with pytest.raises(UrlSecurityError): + check_ssrf_ip("fc00::1") + with pytest.raises(UrlSecurityError): + check_ssrf_ip("fe80::1") + + def test_ipv6_public_passes(self): + """IPv6 公网地址通过.""" + check_ssrf_ip("2001:4860:4860::8888") + + def test_invalid_ip_raises_value_error(self): + """非法 IP 抛出 ValueError(不是 UrlSecurityError).""" + with pytest.raises(ValueError): + check_ssrf_ip("not-an-ip") + with pytest.raises(ValueError): + check_ssrf_ip("999.999.999.999") + + +# ═══════════════════════════════════════════════════════════════════════════════ +# is_ip_address +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestIsIpAddress: + """IP 地址判断测试.""" + + def test_ipv4_true(self): + """IPv4 地址返回 True.""" + assert is_ip_address("127.0.0.1") is True + assert is_ip_address("8.8.8.8") is True + assert is_ip_address("0.0.0.0") is True + + def test_ipv6_true(self): + """IPv6 地址返回 True.""" + assert is_ip_address("::1") is True + assert is_ip_address("2001:db8::1") is True + + def test_hostname_false(self): + """主机名返回 False.""" + assert is_ip_address("example.com") is False + assert is_ip_address("localhost") is False + assert is_ip_address("sub.domain.org") is False + + def test_empty_string_false(self): + """空字符串返回 False.""" + assert is_ip_address("") is False + + def test_invalid_ip_false(self): + """非法 IP 返回 False.""" + assert is_ip_address("999.999.999.999") is False + assert is_ip_address("1234") is False + assert is_ip_address("abc.def") is False + + +# ═══════════════════════════════════════════════════════════════════════════════ +# validate_url_basic +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestValidateUrlBasic: + """URL 基础校验测试.""" + + def test_normal_https_url_passes(self): + """正常 HTTPS URL 通过.""" + result = validate_url_basic("https://example.com/path") + assert result == "https://example.com/path" + + def test_normal_http_url_passes(self): + """正常 HTTP URL 通过.""" + result = validate_url_basic("http://example.com/path") + assert result == "http://example.com/path" + + def test_empty_url_rejected(self): + """空 URL 被拒.""" + with pytest.raises(UrlSecurityError, match="为空"): + validate_url_basic("") + + def test_none_url_not_passed_as_str(self): + """None 作为 URL(这里只测空字符串).""" + # 空字符串被拒 + with pytest.raises(UrlSecurityError): + validate_url_basic("") + + def test_too_long_url_rejected(self): + """超长 URL 被拒.""" + long_url = "https://example.com/" + "a" * 3000 + with pytest.raises(UrlSecurityError, match="过长"): + validate_url_basic(long_url) + + def test_invalid_scheme_rejected(self): + """非法 scheme 被拒.""" + with pytest.raises(UrlSecurityError, match="scheme"): + validate_url_basic("ftp://example.com/file") + with pytest.raises(UrlSecurityError): + validate_url_basic("file:///etc/passwd") + with pytest.raises(UrlSecurityError): + validate_url_basic("javascript:alert(1)") + + def test_missing_hostname_rejected(self): + """缺少主机名被拒.""" + with pytest.raises(UrlSecurityError, match="主机名"): + validate_url_basic("https:///path") + + def test_localhost_rejected(self): + """localhost 被拒.""" + with pytest.raises(UrlSecurityError): + validate_url_basic("https://localhost/api") + + def test_internal_domain_rejected(self): + """内网域名被拒.""" + with pytest.raises(UrlSecurityError): + validate_url_basic("http://server.local/api") + + def test_non_standard_port_rejected(self): + """非标准端口被拒.""" + with pytest.raises(UrlSecurityError, match="端口"): + validate_url_basic("https://example.com:8080/") + with pytest.raises(UrlSecurityError): + validate_url_basic("http://example.com:3000/") + + def test_port_80_ok(self): + """80 端口允许.""" + validate_url_basic("http://example.com:80/path") + + def test_port_443_ok(self): + """443 端口允许.""" + validate_url_basic("https://example.com:443/path") + + def test_no_port_ok(self): + """无端口默认允许.""" + validate_url_basic("https://example.com/path") + + def test_direct_ip_rejected_by_default(self): + """默认禁止直接 IP 访问.""" + with pytest.raises(UrlSecurityError, match="直接 IP"): + validate_url_basic("https://8.8.8.8/path") + + def test_direct_ip_allowed_when_enabled(self): + """allow_direct_ip=True 时允许公网 IP.""" + validate_url_basic("https://8.8.8.8/path", allow_direct_ip=True) + + def test_direct_ip_private_still_blocked(self): + """即使 allow_direct_ip,内网 IP 仍被拒.""" + with pytest.raises(UrlSecurityError, match="内网"): + validate_url_basic("https://192.168.1.1/", allow_direct_ip=True) + + def test_direct_ip_loopback_still_blocked(self): + """回环 IP 即使开启 direct_ip 也被拒.""" + with pytest.raises(UrlSecurityError): + validate_url_basic("https://127.0.0.1/", allow_direct_ip=True) + + def test_trusted_domains_pass(self): + """可信域名列表内的域名通过.""" + trusted = {"example.com", "cdn.com"} + validate_url_basic("https://api.example.com/path", trusted_domains=trusted) + validate_url_basic("https://cdn.com/asset.jpg", trusted_domains=trusted) + + def test_untrusted_domain_rejected(self): + """不在可信域名列表中的域名被拒.""" + trusted = {"example.com"} + with pytest.raises(UrlSecurityError, match="白名单"): + validate_url_basic("https://evil.com/malware", trusted_domains=trusted) + + def test_trusted_domain_subdomain_pass(self): + """可信域名的子域名通过.""" + trusted = {"example.com"} + validate_url_basic("https://sub.example.com/a", trusted_domains=trusted) + validate_url_basic("https://a.b.example.com/b", trusted_domains=trusted) + + def test_return_value_is_original_url(self): + """返回原始 URL 字符串.""" + url = "https://example.com/path?query=value#frag" + assert validate_url_basic(url) == url + + def test_metadata_ip_rejected(self): + """云元数据 IP 被内部主机名检查拦截.""" + with pytest.raises(UrlSecurityError): + validate_url_basic("http://169.254.169.254/latest/meta-data/") + + +# ═══════════════════════════════════════════════════════════════════════════════ +# is_url_basic_safe +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestIsUrlBasicSafe: + """便捷函数 is_url_basic_safe 测试.""" + + def test_safe_url_returns_true(self): + """安全 URL 返回 True.""" + assert is_url_basic_safe("https://example.com/") is True + assert is_url_basic_safe("http://example.org/path") is True + + def test_unsafe_url_returns_false(self): + """不安全 URL 返回 False.""" + assert is_url_basic_safe("https://localhost/") is False + assert is_url_basic_safe("ftp://example.com/") is False + assert is_url_basic_safe("") is False + + def test_trusted_domains_param(self): + """支持 trusted_domains 参数.""" + trusted = {"example.com"} + assert is_url_basic_safe("https://other.com/", trusted_domains=trusted) is False + assert is_url_basic_safe("https://example.com/", trusted_domains=trusted) is True + + def test_allow_direct_ip_param(self): + """支持 allow_direct_ip 参数.""" + assert is_url_basic_safe("https://8.8.8.8/") is False + assert is_url_basic_safe("https://8.8.8.8/", allow_direct_ip=True) is True + + def test_no_exceptions_raised(self): + """不抛出异常,只返回 bool.""" + # 各种边界情况都不抛异常 + try: + is_url_basic_safe("") + is_url_basic_safe("not a url") + is_url_basic_safe("http://" + "a" * 3000) + except UrlSecurityError: + pytest.fail("is_url_basic_safe should not raise UrlSecurityError") + + +# ═══════════════════════════════════════════════════════════════════════════════ +# validate_magic_number +# ═══════════════════════════════════════════════════════════════════════════════ + + +class TestValidateMagicNumber: + """魔数校验测试.""" + + def test_jpeg_valid(self): + """JPEG 文件通过.""" + # JPEG 文件头: FF D8 FF + jpeg_header = b"\xff\xd8\xff\xe0\x00\x10JFIF\x00" + validate_magic_number(jpeg_header, {"image/jpeg"}) + + def test_png_valid(self): + """PNG 文件通过.""" + png_header = b"\x89PNG\r\n\x1a\n\x00\x00\x00" + validate_magic_number(png_header, {"image/png"}) + + def test_gif_valid(self): + """GIF 文件通过(GIF89a 和 GIF87a).""" + validate_magic_number(b"GIF89a...", {"image/gif"}) + validate_magic_number(b"GIF87a...", {"image/gif"}) + + def test_wav_valid(self): + """WAV 文件通过(RIFF + WAVE).""" + wav_header = b"RIFF\x00\x00\x00\x00WAVEfmt " + validate_magic_number(wav_header, {"audio/wav"}) + + def test_mp3_id3_valid(self): + """带 ID3 标签的 MP3 通过.""" + mp3_header = b"ID3\x03\x00\x00\x00\x00\x0f\x76" + validate_magic_number(mp3_header, {"audio/mpeg"}) + + def test_mp3_sync_valid(self): + """不带 ID3 的 MP3(帧同步字)通过.""" + mp3_header = b"\xff\xfb\x90\x00" + b"\x00" * 32 + validate_magic_number(mp3_header, {"audio/mpeg"}) + + def test_ogg_valid(self): + """OGG 文件通过.""" + validate_magic_number(b"OggS\x00\x00...", {"audio/ogg"}) + + def test_flac_valid(self): + """FLAC 文件通过.""" + validate_magic_number(b"fLaC\x00\x00...", {"audio/flac"}) + + def test_webp_valid(self): + """WebP 文件通过(RIFF + WEBP).""" + webp_header = b"RIFF\x00\x00\x00\x00WEBPVP8 " + validate_magic_number(webp_header, {"image/webp"}) + + def test_bmp_valid(self): + """BMP 文件通过.""" + validate_magic_number(b"BM\x00\x00\x00\x00...", {"image/bmp"}) + + def test_mp4_valid(self): + """MP4 文件通过(ftyp 在偏移 4).""" + mp4_header = b"\x00\x00\x00\x20ftypisom\x00\x00\x02\x00" + validate_magic_number(mp4_header, {"video/mp4"}) + + def test_invalid_format_rejected(self): + """不匹配的格式被拒.""" + with pytest.raises(UrlSecurityError, match="魔数"): + validate_magic_number(b"hello world", {"image/jpeg"}) + + def test_empty_bytes_rejected(self): + """空字节被拒.""" + with pytest.raises(UrlSecurityError, match="为空"): + validate_magic_number(b"", {"image/jpeg"}) + + def test_too_short_bytes_rejected(self): + """字节太短不匹配魔数时被拒.""" + with pytest.raises(UrlSecurityError): + validate_magic_number(b"\xff\xd8", {"image/jpeg"}) # 只2字节,不够JPEG魔数 + + def test_multiple_allowed_types(self): + """允许多种格式时任一匹配即通过.""" + jpeg_header = b"\xff\xd8\xff\xe0\x00\x10JFIF\x00" + validate_magic_number(jpeg_header, {"image/jpeg", "image/png", "image/gif"}) + + def test_wrong_type_rejected(self): + """用 PNG 魔数校验 JPEG 类型失败.""" + jpeg_header = b"\xff\xd8\xff\xe0\x00\x10JFIF\x00" + with pytest.raises(UrlSecurityError): + validate_magic_number(jpeg_header, {"image/png"}) + + def test_unknown_mime_skipped(self): + """未知 MIME 类型(无对应魔数)不阻断.""" + # application/octet-stream 没有魔数定义,直接通过 + validate_magic_number(b"random bytes here", {"application/octet-stream"}) + + def test_allowed_image_mime_types_has_common(self): + """图片 MIME 白名单包含常见类型.""" + assert "image/jpeg" in ALLOWED_IMAGE_MIME_TYPES + assert "image/png" in ALLOWED_IMAGE_MIME_TYPES + + def test_allowed_video_mime_types_has_common(self): + """视频 MIME 白名单包含常见类型.""" + assert "video/mp4" in ALLOWED_VIDEO_MIME_TYPES