"""通用 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 是否允许访问(async,DNS 解析不阻塞事件循环)。 校验顺序(纵深防御,Pre-DNS → Post-DNS): 1. URL 解析 2. 协议校验 3. 主机名阻止列表(Pre-DNS,零副作用) 4. 域名白名单/黑名单(Pre-DNS) 5. IP 字面量校验(含非标准格式和嵌入式 IPv4-in-IPv6,Pre-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()