Files
xiaoxia-saas/tests/unit/domain/test_url_security.py
T

568 lines
23 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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