ForcePilot/backend/package/yuxi/utils/net_security.py
Kris 436c09fb90 refactor(ssrf): 优化SSRF策略的主机名匹配逻辑
将主机名拦截规则从直接使用原始配置改为使用预归一化后的数据,提升匹配效率和一致性
2026-06-13 18:54:01 +08:00

321 lines
12 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 防护模块。
提供纵深防御的 URL 安全校验能力,防止 Agent 通过工具访问内网服务。
校验顺序:协议校验 → 主机名阻止(Pre-DNS) → 域名白/黑名单(Pre-DNS) → IP 字面量校验(Pre-DNS) → DNS 解析后 IP 校验(Post-DNS)
"""
from __future__ import annotations
import ipaddress
import re
import socket
import time
from urllib.parse import urlparse
from pydantic import BaseModel, Field, PrivateAttr
from yuxi.utils.logging_config import logger
class SSRFPolicy(BaseModel):
"""SSRF 防护策略配置。"""
allowed_schemes: list[str] = Field(
default_factory=lambda: ["http", "https"],
description="允许的 URL 协议",
)
blocked_ip_ranges: list[str] = Field(
default_factory=lambda: [
"0.0.0.0/8", # 当前网络
"10.0.0.0/8", # 私有网络 A 类
"127.0.0.0/8", # 回环地址
"169.254.0.0/16", # 链路本地(含云元数据 169.254.169.254
"172.16.0.0/12", # 私有网络 B 类
"192.168.0.0/16", # 私有网络 C 类
"198.18.0.0/15", # RFC 2544 基准测试(含 fake-ip 代理)
"100.64.0.0/10", # Carrier-grade NAT
"224.0.0.0/4", # 组播
"240.0.0.0/4", # 保留
"255.255.255.255/32", # 广播
"::/128", # IPv6 未指定
"::1/128", # IPv6 回环
"fc00::/7", # IPv6 唯一本地
"fe80::/10", # IPv6 链路本地
"ff00::/8", # IPv6 组播
"100::/64", # IPv6 丢弃前缀
"2001:2::/48", # IPv6 基准测试
"2001:20::/28", # ORCHIDv2
"fec0::/10", # 废弃的 site-local
],
description="阻止的 IP 地址范围CIDR",
)
allowed_domains: list[str] = Field(
default_factory=list,
description="域名白名单(空=不限制域名,仅限制 IP。支持通配符 *.example.com",
)
blocked_domains: list[str] = Field(
default_factory=list,
description="域名黑名单(精确匹配)",
)
blocked_hostnames: list[str] = Field(
default_factory=lambda: [
"localhost",
"localhost.localdomain",
"metadata.google.internal",
],
description="阻止的主机名精确匹配Pre-DNS 阶段即阻止)",
)
blocked_hostname_suffixes: list[str] = Field(
default_factory=lambda: [".localhost", ".local", ".internal"],
description="阻止的主机名后缀Pre-DNS 阶段即阻止)",
)
dns_pinned_ttl: int = Field(
default=300,
description="DNS 缓存时间(秒),防重绑定",
)
_blocked_networks: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = PrivateAttr(default_factory=list)
_blocked_domains_normalized: set[str] = PrivateAttr(default_factory=set)
_blocked_hostnames_normalized: set[str] = PrivateAttr(default_factory=set)
_blocked_hostname_suffixes_normalized: list[str] = PrivateAttr(default_factory=list)
model_config = {"arbitrary_types_allowed": True}
def model_post_init(self, __context) -> None:
self._blocked_networks = [
ipaddress.ip_network(cidr) for cidr in self.blocked_ip_ranges
]
self._blocked_domains_normalized = {
_normalize_hostname(d) for d in self.blocked_domains
}
self._blocked_hostnames_normalized = {
_normalize_hostname(h) for h in self.blocked_hostnames
}
self._blocked_hostname_suffixes_normalized = [
_normalize_hostname(s) for s in self.blocked_hostname_suffixes
]
class SSRFBlockedError(Exception):
"""SSRF 防护拦截异常。"""
def __init__(self, message: str, hostname: str = ""):
super().__init__(message)
self.hostname = hostname
def _normalize_hostname(hostname: str) -> str:
"""主机名规范化:小写化、去尾点、去 IPv6 方括号。"""
normalized = hostname.strip().lower().rstrip(".")
if normalized.startswith("[") and normalized.endswith("]"):
normalized = normalized[1:-1]
return normalized
def _is_blocked_hostname(hostname: str, policy: SSRFPolicy) -> bool:
"""检查主机名是否在阻止列表中Pre-DNS零副作用"""
normalized = _normalize_hostname(hostname)
if not normalized:
return True # fail-closed
if normalized in policy._blocked_hostnames_normalized:
return True
return any(normalized.endswith(suffix) for suffix in policy._blocked_hostname_suffixes_normalized)
def _matches_domain_pattern(hostname: str, pattern: str) -> bool:
"""检查主机名是否匹配域名模式。
*.example.com 匹配子域名但不匹配 apex 域example.com 本身不匹配 *.example.com
"""
normalized = _normalize_hostname(hostname)
pattern = _normalize_hostname(pattern)
if pattern.startswith("*."):
suffix = pattern[2:]
if not suffix or normalized == suffix:
return False
return normalized.endswith(f".{suffix}")
return normalized == pattern or normalized.endswith(f".{pattern}")
def _is_allowed_by_domain_whitelist(hostname: str, allowed_domains: list[str]) -> bool:
"""检查主机名是否通过域名白名单(空白名单=不限制)。"""
if not allowed_domains:
return True
return any(_matches_domain_pattern(hostname, d) for d in allowed_domains)
def _looks_like_ipv4_literal(hostname: str) -> bool:
"""检测主机名是否看起来像 IPv4 字面量(含非标准格式)。
Python 的 ipaddress 不接受八进制、十六进制、短格式等,但浏览器/HTTP 库可能解析。
"""
parts = hostname.split(".")
if not parts or len(parts) > 4:
return False
if any(not p for p in parts):
return True # 空段,异常格式
return all(
p.isdigit() or (p.lower().startswith("0x") and all(c in "0123456789abcdef" for c in p[2:]))
for p in parts
)
def _is_blocked_ip(addr: ipaddress.IPv4Address | ipaddress.IPv6Address, policy: SSRFPolicy) -> bool:
"""检查 IP 是否在阻止范围内。"""
return any(addr in net for net in policy._blocked_networks)
def _check_ip_literal(hostname: str, policy: SSRFPolicy) -> tuple[bool, str]:
"""检查主机名是否为 IP 字面量(含非标准格式和嵌入式 IPv4-in-IPv6
返回 (is_handled, reason) — is_handled=True 表示已判定is_handled=False 表示不是 IP 字面量。
"""
normalized = _normalize_hostname(hostname)
# 1. 标准 IPv4/IPv6 地址
try:
ip = ipaddress.ip_address(normalized)
if _is_blocked_ip(ip, policy):
return True, f"目标 IP 在阻止范围内: {normalized}"
# IPv4-mapped IPv6: 检查嵌入式 IPv4如 ::ffff:127.0.0.1
if isinstance(ip, ipaddress.IPv6Address) and ip.ipv4_mapped:
if _is_blocked_ip(ip.ipv4_mapped, policy):
return True, f"IPv4-mapped 地址解析到被阻止的 IP: {normalized} -> {ip.ipv4_mapped}"
return True, "" # 合法公网 IP通过
except ValueError:
pass
# 2. 畸形 IPv6 字面量(含冒号但无法解析) — fail-closed
if ":" in normalized:
try:
ipaddress.ip_address(normalized)
except ValueError:
return True, f"畸形 IPv6 地址,按安全策略阻止: {normalized}"
# 3. 非标准 IPv4 字面量(八进制、十六进制、短格式) — fail-closed
if _looks_like_ipv4_literal(normalized):
return True, f"非标准 IPv4 字面量,按安全策略阻止: {normalized}"
return False, "" # 不是 IP 字面量,是域名
def log_ssrf_block(hostname: str, reason: str, context: str = "url-fetch") -> None:
"""记录 SSRF 拦截事件(安全审计日志)。
仅记录主机名,不记录完整 URL避免泄露敏感参数。
"""
logger.warning(f"security: blocked URL fetch ({context}) targetHost={hostname} reason={reason}")
class SSRFValidator:
"""SSRF 防护验证器async 优先,与项目风格一致)。"""
def __init__(self, policy: SSRFPolicy | None = None):
self.policy = policy or SSRFPolicy()
self._dns_cache: dict[str, tuple[list[str], float]] = {}
async def validate_url(self, url: str) -> tuple[bool, str]:
"""校验 URL 是否允许访问asyncDNS 解析不阻塞事件循环)。
校验顺序纵深防御Pre-DNS → Post-DNS
1. URL 解析
2. 协议校验
3. 主机名阻止列表Pre-DNS零副作用
4. 域名白名单/黑名单Pre-DNS
5. IP 字面量校验(含非标准格式和嵌入式 IPv4-in-IPv6Pre-DNS
6. DNS 解析后 IP 校验Post-DNS防重绑定
Returns:
(是否通过, 原因描述)
"""
# 1. 解析 URL
try:
parsed = urlparse(url)
except Exception:
return False, "URL 格式无效"
# 2. 协议校验
if parsed.scheme not in self.policy.allowed_schemes:
return False, f"不允许的协议: {parsed.scheme},仅允许 {self.policy.allowed_schemes}"
hostname = parsed.hostname
if not hostname:
return False, "URL 缺少主机名"
# 3. 主机名阻止列表Pre-DNS
if _is_blocked_hostname(hostname, self.policy):
return False, f"主机名在阻止列表中: {hostname}"
# 4. 域名白名单/黑名单Pre-DNS
if not _is_allowed_by_domain_whitelist(hostname, self.policy.allowed_domains):
return False, f"域名不在白名单中: {hostname}"
if _normalize_hostname(hostname) in self.policy._blocked_domains_normalized:
return False, f"域名在黑名单中: {hostname}"
# 5. IP 字面量校验Pre-DNS
is_handled, reason = _check_ip_literal(hostname, self.policy)
if is_handled:
if reason:
return False, reason
# 合法公网 IP 字面量,通过
return True, "通过"
# 6. DNS 解析后 IP 校验Post-DNS防重绑定
resolved_ips = await self._resolve_with_pin(hostname)
if not resolved_ips:
# DNS 解析失败 — fail-closed
return False, f"DNS 解析失败,按安全策略阻止: {hostname}"
for ip_str in resolved_ips:
try:
ip = ipaddress.ip_address(ip_str)
if _is_blocked_ip(ip, self.policy):
return False, f"DNS 解析到被阻止的 IP: {hostname} -> {ip_str}"
except ValueError:
# 解析出非标准 IP — fail-closed
return False, f"DNS 解析到非标准地址,按安全策略阻止: {hostname} -> {ip_str}"
return True, "通过"
def validate_url_sync(self, url: str) -> tuple[bool, str]:
"""同步版本的 validate_url用于非 async 上下文。DNS 解析会阻塞。"""
import asyncio
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop and loop.is_running():
raise RuntimeError("validate_url_sync 不能在已运行的事件循环中调用,请使用 validate_url")
return asyncio.run(self.validate_url(url))
async def _resolve_with_pin(self, hostname: str) -> list[str]:
"""DNS 解析并缓存,防止重绑定攻击。
在 dns_pinned_ttl 秒内,同一主机名返回缓存的 IP 列表。
"""
import asyncio
now = time.time()
if hostname in self._dns_cache:
cached_ips, cached_at = self._dns_cache[hostname]
if now - cached_at < self.policy.dns_pinned_ttl:
return cached_ips
try:
addr_infos = await asyncio.to_thread(
socket.getaddrinfo, hostname, None, socket.AF_UNSPEC, socket.SOCK_STREAM
)
resolved = list({addr[4][0] for addr in addr_infos})
except (socket.gaierror, OSError):
resolved = [] # fail-closed: 调用方应检查空列表
self._dns_cache[hostname] = (resolved, now)
return resolved
def clear_dns_cache(self) -> None:
"""清除 DNS 缓存(配置变更时调用)。"""
self._dns_cache.clear()