ForcePilot/backend/package/yuxi/utils/net_security.py
Kris bea74470cc feat(utils): 新增加密工具模块并完善SSRF防护能力
1. 新增crypto.py实现基于Fernet的敏感字段加密脱敏通用工具
   - 封装CryptoHelper类,支持密钥注入、环境变量加载
   - 实现单值加解密、递归处理字典列表的敏感字段加解密与脱敏
   - 提供模块级便捷函数和全局实例便于调用
2. 优化net_security.py的SSRF防护逻辑
   - 新增IPv4嵌入IPv6地址检测处理
   - 增加歧义netloc拦截规则
   - 拆分同步异步校验逻辑,新增解析IP返回能力
   - 优化代码结构与注释可读性
2026-06-18 03:23:00 +08:00

461 lines
18 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 asyncio
import ipaddress
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
)
# 已知会在低 32 位嵌入 IPv4 地址的 IPv6 前缀。
# 仅在这些前缀内才提取尾部 IPv4避免对普通公网 IPv6 产生误拦。
_IPV4_EMBEDDED_PREFIXES: tuple[ipaddress.IPv4Network | ipaddress.IPv6Network, ...] = (
ipaddress.ip_network("::/96"), # IPv4-compatible已废弃但部分系统仍解析
ipaddress.ip_network("::ffff:0:0/96"), # IPv4-mapped
ipaddress.ip_network("64:ff9b::/96"), # NAT64 well-known prefix
ipaddress.ip_network("64:ff9b:1::/48"), # NAT64 local-use prefix
)
def _extract_embedded_ipv4(addr: ipaddress.IPv6Address) -> ipaddress.IPv4Address | None:
"""从 IPv6 地址中提取嵌入的 IPv4 地址(如果属于已知嵌入前缀)。"""
if addr.ipv4_mapped is not None:
return addr.ipv4_mapped
for net in _IPV4_EMBEDDED_PREFIXES:
if addr in net:
return ipaddress.IPv4Address(int(addr) & 0xFFFFFFFF)
return None
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 _is_blocked_ip_or_embedded(
addr: ipaddress.IPv4Address | ipaddress.IPv6Address,
policy: SSRFPolicy,
) -> tuple[bool, str]:
"""检查 IP 本身或其嵌入的 IPv4 是否被阻止。
返回 (blocked, reason)。
"""
if _is_blocked_ip(addr, policy):
return True, f"目标 IP 在阻止范围内: {addr}"
if isinstance(addr, ipaddress.IPv6Address):
embedded = _extract_embedded_ipv4(addr)
if embedded is not None and _is_blocked_ip(embedded, policy):
return True, f"IPv6 地址嵌入被阻止的 IPv4: {addr} -> {embedded}"
return False, ""
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)
blocked, reason = _is_blocked_ip_or_embedded(ip, policy)
if blocked:
return True, f"{reason}"
# 包含 IPv4 点分写法但不是已知的 IPv4 嵌入前缀按安全策略阻止fail-closed
if isinstance(ip, ipaddress.IPv6Address) and "." in normalized and _extract_embedded_ipv4(ip) is None:
return True, f"IPv6 地址包含非标准 IPv4 嵌入,按安全策略阻止: {normalized}"
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 _is_ambiguous_netloc(netloc: str) -> bool:
"""检测 netloc 是否包含不同 URL 解析器可能解释不一致的特征。
多个 @ 或反斜杠常被用于 SSRF 解析绕过,按安全策略直接拒绝。
"""
return netloc.count("@") > 1 or "\\" in netloc
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}"
if _is_ambiguous_netloc(parsed.netloc):
return False, "URL 的 netloc 存在歧义,按安全策略阻止"
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)
blocked, reason = _is_blocked_ip_or_embedded(ip, self.policy)
if blocked:
return False, f"DNS 解析到被阻止的 IP: {hostname} -> {ip_str} ({reason})"
except ValueError:
# 解析出非标准 IP — fail-closed
return False, f"DNS 解析到非标准地址,按安全策略阻止: {hostname} -> {ip_str}"
return True, "通过"
async def validate_url_and_resolve(self, url: str) -> tuple[bool, str, list[str]]:
"""校验 URL 并返回缓存的解析 IP 列表。
调用方可以使用返回的 IP 替换主机名发起请求(配合 Host 头),
从而把 DNS 解析结果固定到真实连接上,降低 DNS Rebinding 风险。
Returns:
(是否通过, 原因描述, 解析到的 IP 字符串列表)。未通过时 IP 列表为空。
"""
allowed, reason = await self.validate_url(url)
if not allowed:
return False, reason, []
parsed = urlparse(url)
hostname = parsed.hostname
if not hostname:
return False, "URL 缺少主机名", []
resolved_ips = await self._resolve_with_pin(hostname)
return True, "通过", resolved_ips
def validate_url_sync(self, url: str) -> tuple[bool, str]:
"""同步版本的 validate_url用于非 async 上下文。DNS 解析会阻塞。"""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop and loop.is_running():
raise RuntimeError("validate_url_sync 不能在已运行的事件循环中调用,请使用 validate_url")
return self._validate_url_core_sync(url)
def validate_url_and_resolve_sync(self, url: str) -> tuple[bool, str, list[str]]:
"""同步版本的 validate_url_and_resolve。DNS 解析会阻塞。"""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop and loop.is_running():
raise RuntimeError(
"validate_url_and_resolve_sync 不能在已运行的事件循环中调用,请使用 validate_url_and_resolve"
)
allowed, reason = self._validate_url_core_sync(url)
if not allowed:
return False, reason, []
parsed = urlparse(url)
hostname = parsed.hostname
if not hostname:
return False, "URL 缺少主机名", []
resolved_ips = self._resolve_with_pin_sync(hostname)
return True, "通过", resolved_ips
def _validate_url_core_sync(self, url: str) -> tuple[bool, str]:
"""URL 安全校验核心逻辑(同步版本,供 validate_url_sync 使用)。"""
# 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}"
if _is_ambiguous_netloc(parsed.netloc):
return False, "URL 的 netloc 存在歧义,按安全策略阻止"
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 = self._resolve_with_pin_sync(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)
blocked, reason = _is_blocked_ip_or_embedded(ip, self.policy)
if blocked:
return False, f"DNS 解析到被阻止的 IP: {hostname} -> {ip_str} ({reason})"
except ValueError:
# 解析出非标准 IP — fail-closed
return False, f"DNS 解析到非标准地址,按安全策略阻止: {hostname} -> {ip_str}"
return True, "通过"
async def _resolve_with_pin(self, hostname: str) -> list[str]:
"""DNS 解析并缓存,防止重绑定攻击。
在 dns_pinned_ttl 秒内,同一主机名返回缓存的 IP 列表。
"""
return await asyncio.to_thread(self._resolve_with_pin_sync, hostname)
def _resolve_with_pin_sync(self, hostname: str) -> list[str]:
"""同步版本 DNS 解析与缓存。"""
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 = 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()