ForcePilot/backend/package/yuxi/channel/security/ssrf_guard.py
Kris 849d21120f feat(channel/security): 新增完整安全工具模块
新增一系列安全相关功能:
1.  新增二维码生成工具,支持自定义参数导出图片/Base64/字节流
2.  新增HTML内容安全处理工具,包括标签剥离、转义、URL校验
3.  新增日志敏感信息脱敏工具,支持配置、字典、日志记录脱敏
4.  新增密钥管理运行时,支持从环境变量/文件加载密钥
5.  新增身份链接管理,支持多渠道身份绑定与解析
6.  新增SSRF防护工具,支持域名/IP校验与固定主机
7.  新增安全权限修复工具,修复文件目录权限与配置项
8.  新增外部内容安全处理,支持LLM特殊令牌剥离与注入检测
9.  新增认证限流工具,支持多维度限流与本地回环豁免
10. 新增配对管理工具,支持安全配对码生成与校验
11. 新增白名单管理工具,支持DM/群组/来源白名单校验
2026-05-21 10:32:57 +08:00

212 lines
7.0 KiB
Python

from __future__ import annotations
import asyncio
import ipaddress
import logging
import re
import socket
from typing import Any
logger = logging.getLogger(__name__)
_HOSTNAME_PATTERN = re.compile(
r"^([a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)*[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?$"
)
_BUILTIN_SAFE_HOSTS: tuple[str, ...] = (
"dns.google",
"one.one.one.one",
"doh.opendns.com",
"dns.quad9.net",
)
# RFC 6761 special-use domain names that should be treated as internal
_SPECIAL_USE_DOMAINS = frozenset(
{
"localhost",
"local",
"test",
"invalid",
"example",
"localhost.localdomain",
}
)
_DEFAULT_DENY_BLOCKED_MARKER = "blocked by default-deny"
class SsrfCheckResult:
def __init__(self, safe: bool, reason: str = ""):
self.safe = safe
self.reason = reason
def __repr__(self) -> str:
return f"SsrfCheckResult(safe={self.safe}, reason={self.reason!r})"
class SsrfGuard:
def __init__(self, *, default_deny: bool = True):
self._hostname_allowlist: list[str] = list(_BUILTIN_SAFE_HOSTS)
self._pinned_hosts: dict[str, str] = {}
self._default_deny = default_deny
@property
def default_deny(self) -> bool:
return self._default_deny
@property
def allowlist(self) -> list[str]:
return list(self._hostname_allowlist)
def set_allowlist(self, patterns: list[str]) -> None:
self._hostname_allowlist = patterns
def pin_host(self, hostname: str, ip_address: str) -> None:
self._pinned_hosts[hostname] = ip_address
def check_hostname(self, hostname: str) -> SsrfCheckResult:
if not self._is_valid_hostname(hostname):
return SsrfCheckResult(safe=False, reason=f"Invalid hostname: {hostname}")
if self._is_special_use_domain(hostname):
return SsrfCheckResult(safe=False, reason=f"Special-use domain blocked: {hostname}")
if hostname in self._pinned_hosts:
return SsrfCheckResult(safe=True, reason="pinned host")
if self._is_allowed_hostname(hostname):
return SsrfCheckResult(safe=False, reason="allowlist matched, DNS resolution and IP check required")
if self._default_deny:
return SsrfCheckResult(safe=False, reason=_DEFAULT_DENY_BLOCKED_MARKER)
return SsrfCheckResult(safe=False, reason="hostname requires DNS resolution and IP check")
def check_ip(self, ip_str: str) -> SsrfCheckResult:
try:
ip = ipaddress.ip_address(ip_str)
except ValueError:
return SsrfCheckResult(safe=False, reason=f"Invalid IP address: {ip_str}")
if self._is_private_ip(ip):
return SsrfCheckResult(safe=False, reason=f"Private/internal IP blocked: {ip_str}")
return SsrfCheckResult(safe=True, reason="public IP")
def check_pinned_lookup(self, hostname: str, resolved_ip: str) -> SsrfCheckResult:
pinned = self._pinned_hosts.get(hostname)
if pinned is None:
return SsrfCheckResult(safe=True, reason="no pinned lookup configured")
if pinned != resolved_ip:
return SsrfCheckResult(
safe=False,
reason=f"DNS pinned mismatch for {hostname}: expected {pinned}, got {resolved_ip}",
)
return SsrfCheckResult(safe=True, reason="pinned lookup matched")
async def check_url(self, hostname: str) -> SsrfCheckResult:
result = self.check_hostname(hostname)
if result.safe:
return result
if result.reason == _DEFAULT_DENY_BLOCKED_MARKER:
return SsrfCheckResult(safe=False, reason=f"hostname not in allowlist: {hostname}")
try:
ip_str = await self._resolve_hostname(hostname)
except Exception:
return SsrfCheckResult(safe=False, reason=f"DNS resolution failed for: {hostname}")
result = self.check_ip(ip_str)
if not result.safe:
return result
return self.check_pinned_lookup(hostname, ip_str)
async def check_url_safe(self, hostname: str) -> bool:
result = await self.check_url(hostname)
return result.safe
@classmethod
def from_config(cls, config: dict[str, Any]) -> SsrfGuard:
ssrf_config = config.get("security", {}).get("ssrf", {})
if not isinstance(ssrf_config, dict):
ssrf_config = {}
default_deny = ssrf_config.get("defaultDeny", True)
if not isinstance(default_deny, bool):
default_deny = True
guard = cls(default_deny=default_deny)
allowlist = ssrf_config.get("allowlist", [])
if isinstance(allowlist, list) and allowlist:
merged = list(_BUILTIN_SAFE_HOSTS)
for entry in allowlist:
if isinstance(entry, str) and entry.strip() and entry.strip() not in merged:
merged.append(entry.strip())
guard._hostname_allowlist = merged
pinned = ssrf_config.get("pinned", {})
if isinstance(pinned, dict):
for host, ip in pinned.items():
if isinstance(host, str) and isinstance(ip, str):
guard.pin_host(host.strip(), ip.strip())
return guard
@staticmethod
async def _resolve_hostname(hostname: str) -> str:
loop = asyncio.get_running_loop()
addrinfo = await loop.getaddrinfo(hostname, None, family=socket.AF_UNSPEC, type=socket.SOCK_STREAM)
for family, _, _, _, sockaddr in addrinfo:
if family in (socket.AF_INET, socket.AF_INET6):
return sockaddr[0]
raise OSError(f"No suitable address found for {hostname}")
def _is_valid_hostname(self, hostname: str) -> bool:
if len(hostname) > 253:
return False
if not _HOSTNAME_PATTERN.match(hostname):
return False
return True
def _is_special_use_domain(self, hostname: str) -> bool:
hostname_lower = hostname.lower()
if hostname_lower in _SPECIAL_USE_DOMAINS:
return True
for domain in _SPECIAL_USE_DOMAINS:
if hostname_lower.endswith(f".{domain}"):
return True
return False
def _is_allowed_hostname(self, hostname: str) -> bool:
if not self._hostname_allowlist:
return False
hostname_lower = hostname.lower()
for pattern in self._hostname_allowlist:
pattern_lower = pattern.lower()
if pattern_lower.startswith("*."):
suffix = pattern_lower[1:]
if hostname_lower.endswith(suffix):
return True
elif pattern_lower == hostname_lower:
return True
return False
@staticmethod
def _is_private_ip(ip: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool:
return (
ip.is_loopback
or ip.is_private
or ip.is_link_local
or ip.is_multicast
or ip.is_reserved
or ip.is_unspecified
)
ssrf_guard = SsrfGuard()