ForcePilot/backend/test/unit/utils/test_net_security.py
Kris 9ce17c1662 feat: 新增通用 SSRF 防护模块并重构 URL 校验逻辑
- 新增完整的 SSRF 防护工具类与配置项
- 重构 URL 校验逻辑,替换原有白名单校验实现
- 新增 DNS 缓存与防重绑定防护
- 兼容原有 YUXI_URL_WHITELIST 环境变量配置
- 新增配套单元测试覆盖防护逻辑
2026-06-13 18:22:38 +08:00

353 lines
13 KiB
Python
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.

"""SSRF 防护模块单元测试。"""
import ipaddress
import pytest
from unittest.mock import patch, AsyncMock
from yuxi.utils.net_security import (
SSRFPolicy,
SSRFValidator,
SSRFBlockedError,
_normalize_hostname,
_is_blocked_hostname,
_matches_domain_pattern,
_is_allowed_by_domain_whitelist,
_looks_like_ipv4_literal,
_check_ip_literal,
)
class TestNormalizeHostname:
def test_lowercase(self):
assert _normalize_hostname("Example.COM") == "example.com"
def test_strip_trailing_dot(self):
assert _normalize_hostname("example.com.") == "example.com"
def test_strip_ipv6_brackets(self):
assert _normalize_hostname("[::1]") == "::1"
def test_empty(self):
assert _normalize_hostname("") == ""
def test_whitespace(self):
assert _normalize_hostname(" example.com ") == "example.com"
class TestBlockedHostname:
def test_localhost(self):
assert _is_blocked_hostname("localhost", SSRFPolicy())
def test_localhost_localdomain(self):
assert _is_blocked_hostname("localhost.localdomain", SSRFPolicy())
def test_metadata_google_internal(self):
assert _is_blocked_hostname("metadata.google.internal", SSRFPolicy())
def test_dot_localhost_suffix(self):
assert _is_blocked_hostname("sub.localhost", SSRFPolicy())
def test_dot_local_suffix(self):
assert _is_blocked_hostname("printer.local", SSRFPolicy())
def test_dot_internal_suffix(self):
assert _is_blocked_hostname("service.internal", SSRFPolicy())
def test_normal_hostname_not_blocked(self):
assert not _is_blocked_hostname("example.com", SSRFPolicy())
def test_empty_hostname_blocked(self):
assert _is_blocked_hostname("", SSRFPolicy()) # fail-closed
def test_case_insensitive(self):
assert _is_blocked_hostname("Localhost", SSRFPolicy())
assert _is_blocked_hostname("SERVICE.INTERNAL", SSRFPolicy())
class TestDomainPattern:
def test_exact_match(self):
assert _matches_domain_pattern("api.example.com", "api.example.com")
def test_subdomain_match(self):
assert _matches_domain_pattern("sub.api.example.com", "api.example.com")
def test_wildcard_match_subdomain(self):
assert _matches_domain_pattern("sub.example.com", "*.example.com")
def test_wildcard_no_match_apex(self):
# *.example.com 不匹配 example.com 本身
assert not _matches_domain_pattern("example.com", "*.example.com")
def test_wildcard_match_deep_subdomain(self):
assert _matches_domain_pattern("a.b.example.com", "*.example.com")
def test_case_insensitive(self):
assert _matches_domain_pattern("Sub.Example.COM", "*.example.com")
class TestDomainWhitelist:
def test_empty_whitelist_allows_all(self):
assert _is_allowed_by_domain_whitelist("any.com", [])
def test_whitelist_match(self):
assert _is_allowed_by_domain_whitelist("api.example.com", ["api.example.com"])
def test_whitelist_no_match(self):
assert not _is_allowed_by_domain_whitelist("other.com", ["api.example.com"])
def test_whitelist_wildcard(self):
assert _is_allowed_by_domain_whitelist("sub.example.com", ["*.example.com"])
class TestIpv4LiteralDetection:
def test_octal(self):
assert _looks_like_ipv4_literal("0177.0.0.1")
def test_hex(self):
assert _looks_like_ipv4_literal("0x7f.0.0.1")
def test_short_format(self):
assert _looks_like_ipv4_literal("127.1")
def test_decimal_only(self):
assert _looks_like_ipv4_literal("2130706433")
def test_normal_hostname_not_detected(self):
assert not _looks_like_ipv4_literal("example.com")
def test_standard_ipv4_detected(self):
# 标准 IPv4 也看起来像字面量,但会被标准解析优先处理
assert _looks_like_ipv4_literal("127.0.0.1")
def test_too_many_parts(self):
assert not _looks_like_ipv4_literal("1.2.3.4.5")
def test_empty_part(self):
assert _looks_like_ipv4_literal("1..3.4")
class TestIpLiteralCheck:
def test_blocked_private_ipv4(self):
handled, reason = _check_ip_literal("192.168.1.1", SSRFPolicy())
assert handled
assert reason
def test_blocked_loopback(self):
handled, reason = _check_ip_literal("127.0.0.1", SSRFPolicy())
assert handled
assert reason
def test_blocked_cloud_metadata(self):
handled, reason = _check_ip_literal("169.254.169.254", SSRFPolicy())
assert handled
assert reason
def test_blocked_ipv4_mapped_ipv6(self):
handled, reason = _check_ip_literal("::ffff:127.0.0.1", SSRFPolicy())
assert handled
assert reason
def test_blocked_ipv6_loopback(self):
handled, reason = _check_ip_literal("::1", SSRFPolicy())
assert handled
assert reason
def test_allowed_public_ipv4(self):
handled, reason = _check_ip_literal("8.8.8.8", SSRFPolicy())
assert handled
assert not reason
def test_nonstandard_ipv4_blocked(self):
handled, reason = _check_ip_literal("0177.0.0.1", SSRFPolicy())
assert handled
assert reason # fail-closed
def test_domain_not_handled(self):
handled, reason = _check_ip_literal("example.com", SSRFPolicy())
assert not handled
def test_blocked_10_network(self):
handled, reason = _check_ip_literal("10.0.0.1", SSRFPolicy())
assert handled
assert reason
def test_blocked_172_network(self):
handled, reason = _check_ip_literal("172.16.0.1", SSRFPolicy())
assert handled
assert reason
def test_blocked_zero_network(self):
handled, reason = _check_ip_literal("0.0.0.0", SSRFPolicy())
assert handled
assert reason
def test_allowed_public_ipv6(self):
handled, reason = _check_ip_literal("2001:4860:4860::8888", SSRFPolicy())
assert handled
assert not reason
class TestSSRFValidator:
def setup_method(self):
self.validator = SSRFValidator()
@pytest.mark.parametrize("url", [
"http://127.0.0.1:8080/admin",
"http://192.168.1.1/config",
"http://10.0.0.1/internal",
"http://172.16.0.1/secret",
"http://169.254.169.254/latest/meta-data/",
"http://0.0.0.0/",
"ftp://example.com/file",
"http://metadata.google.internal/computeMetadata/v1/",
"http://service.internal/api",
"http://printer.local/ipp",
])
@pytest.mark.asyncio
async def test_blocked_urls(self, url):
allowed, reason = await self.validator.validate_url(url)
assert not allowed, f"应阻止 {url},但通过了校验"
@pytest.mark.asyncio
async def test_blocked_localhost(self):
allowed, reason = await self.validator.validate_url("http://localhost:5432/db")
assert not allowed
@pytest.mark.parametrize("url", [
"https://api.openai.com/v1/chat/completions",
"https://www.baidu.com/search",
"http://example.com/page",
])
@pytest.mark.asyncio
async def test_allowed_urls(self, url):
validator = SSRFValidator()
with patch.object(validator, "_resolve_with_pin", new_callable=AsyncMock, return_value=["1.2.3.4"]):
allowed, reason = await validator.validate_url(url)
assert allowed, f"应允许 {url},但被阻止: {reason}"
@pytest.mark.asyncio
async def test_domain_whitelist(self):
policy = SSRFPolicy(allowed_domains=["api.example.com", "trusted.org"])
validator = SSRFValidator(policy)
with patch.object(validator, "_resolve_with_pin", new_callable=AsyncMock, return_value=["1.2.3.4"]):
allowed, _ = await validator.validate_url("https://api.example.com/endpoint")
assert allowed
allowed, _ = await validator.validate_url("https://other.com/endpoint")
assert not allowed
@pytest.mark.asyncio
async def test_domain_whitelist_wildcard(self):
policy = SSRFPolicy(allowed_domains=["*.example.com"])
validator = SSRFValidator(policy)
with patch.object(validator, "_resolve_with_pin", new_callable=AsyncMock, return_value=["1.2.3.4"]):
allowed, _ = await validator.validate_url("https://sub.example.com/endpoint")
assert allowed
# *.example.com 不匹配 apex 域
allowed, _ = await validator.validate_url("https://example.com/endpoint")
assert not allowed
@pytest.mark.asyncio
async def test_domain_blacklist(self):
policy = SSRFPolicy(blocked_domains=["evil.com"])
validator = SSRFValidator(policy)
allowed, _ = await validator.validate_url("https://evil.com/attack")
assert not allowed
with patch.object(validator, "_resolve_with_pin", new_callable=AsyncMock, return_value=["1.2.3.4"]):
allowed, _ = await validator.validate_url("https://safe.com/page")
assert allowed
@pytest.mark.asyncio
async def test_dns_resolution_blocked_private(self):
"""DNS 解析到私有 IP 时应阻止"""
validator = SSRFValidator()
with patch.object(validator, "_resolve_with_pin", new_callable=AsyncMock, return_value=["10.0.0.1"]):
allowed, reason = await validator.validate_url("https://looks-legit.com/api")
assert not allowed
assert "DNS 解析到被阻止的 IP" in reason
@pytest.mark.asyncio
async def test_dns_resolution_fail_closed(self):
"""DNS 解析失败时应阻止fail-closed"""
validator = SSRFValidator()
with patch.object(validator, "_resolve_with_pin", new_callable=AsyncMock, return_value=[]):
allowed, reason = await validator.validate_url("https://unresolvable.example.com/api")
assert not allowed
assert "DNS 解析失败" in reason
@pytest.mark.asyncio
async def test_ipv4_mapped_ipv6_blocked(self):
"""IPv4-mapped IPv6 地址应提取并校验嵌入式 IPv4"""
allowed, reason = await self.validator.validate_url("http://[::ffff:127.0.0.1]/admin")
assert not allowed
@pytest.mark.asyncio
async def test_nonstandard_ipv4_literal_blocked(self):
"""非标准 IPv4 字面量应被阻止fail-closed"""
handled, reason = _check_ip_literal("0177.0.0.1", SSRFPolicy())
assert handled
assert reason
@pytest.mark.asyncio
async def test_empty_url(self):
allowed, reason = await self.validator.validate_url("")
assert not allowed
@pytest.mark.asyncio
async def test_no_hostname(self):
allowed, reason = await self.validator.validate_url("http:///path")
assert not allowed
@pytest.mark.asyncio
async def test_invalid_scheme(self):
allowed, reason = await self.validator.validate_url("ftp://example.com/file")
assert not allowed
assert "协议" in reason
@pytest.mark.asyncio
async def test_dns_pinning(self):
"""DNS Pinning 在 TTL 内返回缓存结果"""
validator = SSRFValidator(SSRFPolicy(dns_pinned_ttl=60))
# 首次解析
with patch.object(validator, "_resolve_with_pin", new_callable=AsyncMock, return_value=["1.2.3.4"]) as mock:
await validator.validate_url("https://example.com/page")
assert mock.call_count == 1
class TestSSRFPolicyModel:
def test_default_policy(self):
policy = SSRFPolicy()
assert policy.allowed_schemes == ["http", "https"]
assert policy.dns_pinned_ttl == 300
assert len(policy._blocked_networks) > 0
def test_custom_policy(self):
policy = SSRFPolicy(
allowed_domains=["api.example.com"],
blocked_domains=["evil.com"],
dns_pinned_ttl=60,
)
assert policy.allowed_domains == ["api.example.com"]
assert policy.blocked_domains == ["evil.com"]
assert policy.dns_pinned_ttl == 60
def test_blocked_networks_initialized(self):
policy = SSRFPolicy()
assert ipaddress.ip_network("10.0.0.0/8") in policy._blocked_networks
assert ipaddress.ip_network("127.0.0.0/8") in policy._blocked_networks
assert ipaddress.ip_network("::1/128") in policy._blocked_networks
def test_blocked_networks_includes_cloud_metadata(self):
policy = SSRFPolicy()
# 验证 169.254.169.254 在某个阻止网络内
assert any(ipaddress.ip_address("169.254.169.254") in net for net in policy._blocked_networks)
class TestSSRFBlockedError:
def test_error_message(self):
error = SSRFBlockedError("blocked", hostname="evil.com")
assert str(error) == "blocked"
assert error.hostname == "evil.com"