"""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