feat(utils): 新增加密工具模块并完善SSRF防护能力
1. 新增crypto.py实现基于Fernet的敏感字段加密脱敏通用工具 - 封装CryptoHelper类,支持密钥注入、环境变量加载 - 实现单值加解密、递归处理字典列表的敏感字段加解密与脱敏 - 提供模块级便捷函数和全局实例便于调用 2. 优化net_security.py的SSRF防护逻辑 - 新增IPv4嵌入IPv6地址检测处理 - 增加歧义netloc拦截规则 - 拆分同步异步校验逻辑,新增解析IP返回能力 - 优化代码结构与注释可读性
This commit is contained in:
parent
32b213593e
commit
bea74470cc
348
backend/package/yuxi/utils/crypto.py
Normal file
348
backend/package/yuxi/utils/crypto.py
Normal file
@ -0,0 +1,348 @@
|
||||
"""敏感字段加密与脱敏通用工具。
|
||||
|
||||
基于 Fernet 对称加密,支持递归处理配置字典中的敏感字段。
|
||||
所有加密、解密、脱敏、密钥校验逻辑均封装在 ``CryptoHelper`` 类中,
|
||||
模块级便捷函数仅作为默认实例的薄封装,便于全局调用与测试注入。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Literal
|
||||
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
|
||||
from yuxi.utils import logger
|
||||
|
||||
|
||||
class CryptoHelper:
|
||||
"""敏感字段加密与脱敏工具(完整闭环)。
|
||||
|
||||
基于 Fernet 对称加密,支持递归处理配置字典中的敏感字段。
|
||||
|
||||
密钥来源优先级:
|
||||
1. 构造函数注入的 ``fernet`` 实例
|
||||
2. 构造函数注入的 ``key`` 字符串
|
||||
3. 环境变量 ``ENCRYPTION_KEY``
|
||||
|
||||
未提供密钥时,加密/解密操作将透传原值(开发环境兼容),并在首次
|
||||
访问时记录一次警告日志;生产环境可通过 ``validate_for_environment``
|
||||
强制校验密钥存在性。
|
||||
|
||||
敏感字段判定:字段名(小写归一化后)命中 ``_SENSITIVE_KEYWORDS`` 中
|
||||
任一关键词即视为敏感。设计原则是宁可漏判不可误判,因此避免使用过于
|
||||
宽泛的词(如 "key" / "cert" / "auth" / "id"),防止误伤
|
||||
``auth_type`` / ``auth_config`` / ``content-type`` / ``client_id``
|
||||
等非敏感字段。
|
||||
"""
|
||||
|
||||
# 加密密钥环境变量名
|
||||
_ENV_KEY_NAME = "ENCRYPTION_KEY"
|
||||
|
||||
# 用于识别 Fernet 加密值的统一前缀
|
||||
_ENCRYPTION_PREFIX = "enc:"
|
||||
|
||||
# 敏感字段关键词:字段名(小写归一化后)命中任一关键词即视为敏感。
|
||||
# 设计原则:宁可漏判(漏判可由调用方补充),不可误判(误判会破坏正常字段)。
|
||||
# 因此避免使用过于宽泛的词(如 "key" / "cert" / "auth" / "id"),
|
||||
# 防止误伤 auth_type / auth_config / content-type / client_id 等非敏感字段。
|
||||
_SENSITIVE_KEYWORDS = (
|
||||
# 密码类
|
||||
"password",
|
||||
"passwd",
|
||||
"pwd",
|
||||
# 令牌类
|
||||
"token",
|
||||
"access_token",
|
||||
"refresh_token",
|
||||
"id_token",
|
||||
"bearer",
|
||||
"authorization",
|
||||
# 密钥类
|
||||
"secret",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"api_secret",
|
||||
"access_key",
|
||||
"secret_key",
|
||||
"private_key",
|
||||
"client_secret",
|
||||
"client_key",
|
||||
# 证书类
|
||||
"ca_cert",
|
||||
"client_cert",
|
||||
"certificate",
|
||||
# 凭证类
|
||||
"credential",
|
||||
"credentials",
|
||||
)
|
||||
|
||||
def __init__(self, key: str | None = None, fernet: Fernet | None = None) -> None:
|
||||
"""初始化加密助手。
|
||||
|
||||
Args:
|
||||
key: Fernet 密钥字符串。提供时优先使用,忽略环境变量。
|
||||
fernet: 已构造的 Fernet 实例。提供时优先于 ``key``。
|
||||
"""
|
||||
if fernet is not None:
|
||||
self._fernet: Fernet | None = fernet
|
||||
self._injected_key: str | None = None
|
||||
elif key is not None:
|
||||
self._fernet = Fernet(key.encode())
|
||||
self._injected_key = key
|
||||
else:
|
||||
self._fernet = None
|
||||
self._injected_key = None
|
||||
self._missing_logged = False
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 密钥管理
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> CryptoHelper:
|
||||
"""从环境变量 ``ENCRYPTION_KEY`` 创建实例。
|
||||
|
||||
环境变量未设置时返回未配置密钥的实例(开发环境兼容)。
|
||||
"""
|
||||
key = os.getenv(cls._ENV_KEY_NAME)
|
||||
return cls(key=key) if key else cls()
|
||||
|
||||
def has_key(self) -> bool:
|
||||
"""检查是否存在有效密钥(不触发日志,不抛异常)。"""
|
||||
if self._fernet is not None:
|
||||
return True
|
||||
if self._injected_key is not None:
|
||||
return True
|
||||
return bool(os.getenv(self._ENV_KEY_NAME))
|
||||
|
||||
def _get_fernet(self) -> Fernet | None:
|
||||
"""获取 Fernet 实例,按密钥来源优先级解析。
|
||||
|
||||
密钥未配置时记录一次警告并返回 None;密钥格式无效时抛出 ValueError。
|
||||
"""
|
||||
if self._fernet is not None:
|
||||
return self._fernet
|
||||
|
||||
key = self._injected_key if self._injected_key is not None else os.getenv(self._ENV_KEY_NAME)
|
||||
if not key:
|
||||
if not self._missing_logged:
|
||||
logger.warning(
|
||||
f"环境变量 {self._ENV_KEY_NAME} 未设置,配置中的敏感字段将以明文存储。"
|
||||
"生产环境请使用 Fernet.generate_key() 生成密钥并配置。"
|
||||
)
|
||||
self._missing_logged = True
|
||||
return None
|
||||
|
||||
try:
|
||||
self._fernet = Fernet(key.encode())
|
||||
except Exception as exc:
|
||||
raise ValueError(f"{self._ENV_KEY_NAME} 格式无效,无法初始化 Fernet") from exc
|
||||
return self._fernet
|
||||
|
||||
def validate_for_environment(self, environment: str) -> None:
|
||||
"""校验生产环境是否正确配置加密密钥。
|
||||
|
||||
非生产环境跳过校验;生产环境要求密钥必须存在且能初始化 Fernet,
|
||||
否则抛出 ValueError。
|
||||
"""
|
||||
if environment != "production":
|
||||
return
|
||||
|
||||
if not self.has_key():
|
||||
raise ValueError(f"生产环境必须设置 {self._ENV_KEY_NAME}")
|
||||
|
||||
try:
|
||||
self._get_fernet()
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"{self._ENV_KEY_NAME} 格式无效") from exc
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 单值加密/解密
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def encrypt(self, value: str) -> str:
|
||||
"""加密字符串,返回带 ``enc:`` 前缀的密文。
|
||||
|
||||
密钥未配置时透传原值。
|
||||
"""
|
||||
fernet = self._get_fernet()
|
||||
if fernet is None:
|
||||
return value
|
||||
encrypted = fernet.encrypt(value.encode())
|
||||
return f"{self._ENCRYPTION_PREFIX}{encrypted.decode()}"
|
||||
|
||||
def decrypt(self, value: str) -> str:
|
||||
"""解密带 ``enc:`` 前缀的密文,返回明文。
|
||||
|
||||
非 ``enc:`` 前缀的值视为未加密,透传返回;密钥未配置时记录警告
|
||||
并透传;解密失败(密钥变更或数据损坏)抛出 ValueError。
|
||||
"""
|
||||
if not isinstance(value, str) or not value.startswith(self._ENCRYPTION_PREFIX):
|
||||
return value
|
||||
|
||||
fernet = self._get_fernet()
|
||||
if fernet is None:
|
||||
logger.warning("缺少加密密钥,无法解密配置中的敏感字段")
|
||||
return value
|
||||
|
||||
try:
|
||||
decrypted = fernet.decrypt(value[len(self._ENCRYPTION_PREFIX) :].encode())
|
||||
except InvalidToken as exc:
|
||||
raise ValueError("配置中的敏感字段解密失败,密钥可能已变更") from exc
|
||||
return decrypted.decode()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 递归加密/解密/脱敏
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def encrypt_sensitive_fields(self, config: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""递归加密配置字典中的敏感字符串字段。仅加密敏感 key 对应的字符串。"""
|
||||
return self._transform_config(config, "encrypt")
|
||||
|
||||
def decrypt_sensitive_fields(self, config: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""递归解密配置字典中的敏感字符串字段。"""
|
||||
return self._transform_config(config, "decrypt")
|
||||
|
||||
def mask_sensitive_fields(self, config: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""返回配置字典的脱敏副本,用于 API 响应。"""
|
||||
result = self._transform_config(config, "mask")
|
||||
return result if result is not None else {}
|
||||
|
||||
def mask_all_string_values(self, obj: Any) -> Any:
|
||||
"""递归将对象中的所有字符串值替换为脱敏占位符。"""
|
||||
if isinstance(obj, dict):
|
||||
return {k: self.mask_all_string_values(v) for k, v in obj.items()}
|
||||
if isinstance(obj, list):
|
||||
return [self.mask_all_string_values(item) for item in obj]
|
||||
if isinstance(obj, str):
|
||||
return "***"
|
||||
return obj
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 内部实现
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@classmethod
|
||||
def _is_sensitive_key(cls, key: str) -> bool:
|
||||
"""判断字段名是否属于敏感字段。"""
|
||||
normalized = re.sub(r"[^a-z0-9_]+", "_", key.lower())
|
||||
return any(keyword in normalized for keyword in cls._SENSITIVE_KEYWORDS)
|
||||
|
||||
def _transform_config(
|
||||
self,
|
||||
config: dict[str, Any] | None,
|
||||
mode: Literal["encrypt", "decrypt", "mask"],
|
||||
) -> dict[str, Any] | None:
|
||||
"""通用配置转换:递归加密/解密/脱敏敏感字段。"""
|
||||
if config is None:
|
||||
return None
|
||||
|
||||
if mode == "encrypt":
|
||||
return self._encrypt_nested(config)
|
||||
if mode == "decrypt":
|
||||
return self._decrypt_nested(config)
|
||||
# mask
|
||||
return self._mask_nested(config)
|
||||
|
||||
def _encrypt_nested(self, obj: Any) -> Any:
|
||||
if isinstance(obj, dict):
|
||||
return {
|
||||
k: self.encrypt(v) if self._is_sensitive_key(k) and isinstance(v, str) else self._encrypt_nested(v)
|
||||
for k, v in obj.items()
|
||||
}
|
||||
if isinstance(obj, list):
|
||||
return [self._encrypt_nested(item) for item in obj]
|
||||
return obj
|
||||
|
||||
def _decrypt_nested(self, obj: Any) -> Any:
|
||||
if isinstance(obj, dict):
|
||||
return {
|
||||
k: self.decrypt(v) if self._is_sensitive_key(k) and isinstance(v, str) else self._decrypt_nested(v)
|
||||
for k, v in obj.items()
|
||||
}
|
||||
if isinstance(obj, list):
|
||||
return [self._decrypt_nested(item) for item in obj]
|
||||
return obj
|
||||
|
||||
def _mask_nested(self, obj: Any) -> Any:
|
||||
if isinstance(obj, dict):
|
||||
return {
|
||||
k: self._mask_string(v) if self._is_sensitive_key(k) and isinstance(v, str) else self._mask_nested(v)
|
||||
for k, v in obj.items()
|
||||
}
|
||||
if isinstance(obj, list):
|
||||
return [self._mask_nested(item) for item in obj]
|
||||
return obj
|
||||
|
||||
@staticmethod
|
||||
def _mask_string(value: str) -> str:
|
||||
"""对字符串进行脱敏,返回统一占位符。"""
|
||||
return "***"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# 模块级默认实例与便捷函数(向后兼容现有调用)
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
_default_helper = CryptoHelper()
|
||||
|
||||
|
||||
def set_crypto_helper(helper: CryptoHelper) -> None:
|
||||
"""替换全局默认 crypto helper,主要用于测试注入。"""
|
||||
global _default_helper
|
||||
_default_helper = helper
|
||||
|
||||
|
||||
def reset_crypto_helper() -> None:
|
||||
"""重置全局默认 crypto helper 为未配置密钥的实例,主要用于测试隔离。"""
|
||||
global _default_helper
|
||||
_default_helper = CryptoHelper()
|
||||
|
||||
|
||||
def get_crypto_helper() -> CryptoHelper:
|
||||
"""获取全局默认 crypto helper 实例。"""
|
||||
return _default_helper
|
||||
|
||||
|
||||
def validate_encryption_key(environment: str) -> None:
|
||||
"""校验生产环境是否正确配置加密密钥(委托给默认实例)。"""
|
||||
_default_helper.validate_for_environment(environment)
|
||||
|
||||
|
||||
def encrypt_sensitive_fields(config: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""递归加密配置字典中的敏感字符串字段(委托给默认实例)。"""
|
||||
return _default_helper.encrypt_sensitive_fields(config)
|
||||
|
||||
|
||||
def decrypt_sensitive_fields(config: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""递归解密配置字典中的敏感字符串字段(委托给默认实例)。"""
|
||||
return _default_helper.decrypt_sensitive_fields(config)
|
||||
|
||||
|
||||
def encrypt(value: str) -> str:
|
||||
"""加密单个字符串,返回带 ``enc:`` 前缀的密文(委托给默认实例)。
|
||||
|
||||
密钥未配置时透传原值。注意:加密不幂等,调用方必须保证传入的是明文。
|
||||
"""
|
||||
return _default_helper.encrypt(value)
|
||||
|
||||
|
||||
def decrypt(value: str) -> str:
|
||||
"""解密单个字符串(委托给默认实例)。
|
||||
|
||||
非 ``enc:`` 前缀的值视为未加密,透传返回;密钥未配置时记录警告并透传;
|
||||
解密失败(密钥变更或数据损坏)抛出 ``ValueError``。
|
||||
"""
|
||||
return _default_helper.decrypt(value)
|
||||
|
||||
|
||||
def mask_sensitive_fields(config: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""返回配置字典的脱敏副本(委托给默认实例)。"""
|
||||
return _default_helper.mask_sensitive_fields(config)
|
||||
|
||||
|
||||
def mask_all_string_values(obj: Any) -> Any:
|
||||
"""递归将对象中的所有字符串值替换为脱敏占位符(委托给默认实例)。"""
|
||||
return _default_helper.mask_all_string_values(obj)
|
||||
@ -1,13 +1,14 @@
|
||||
"""通用 SSRF 防护模块。
|
||||
|
||||
提供纵深防御的 URL 安全校验能力,防止 Agent 通过工具访问内网服务。
|
||||
校验顺序:协议校验 → 主机名阻止(Pre-DNS) → 域名白/黑名单(Pre-DNS) → IP 字面量校验(Pre-DNS) → DNS 解析后 IP 校验(Post-DNS)
|
||||
校验顺序:协议校验 → 主机名阻止(Pre-DNS) → 域名白/黑名单(Pre-DNS) →
|
||||
IP 字面量校验(Pre-DNS) → DNS 解析后 IP 校验(Post-DNS)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import re
|
||||
import socket
|
||||
import time
|
||||
from urllib.parse import urlparse
|
||||
@ -82,18 +83,10 @@ class SSRFPolicy(BaseModel):
|
||||
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
|
||||
]
|
||||
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):
|
||||
@ -155,16 +148,52 @@ def _looks_like_ipv4_literal(hostname: str) -> bool:
|
||||
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
|
||||
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)。
|
||||
|
||||
@ -175,13 +204,13 @@ def _check_ip_literal(hostname: str, policy: SSRFPolicy) -> tuple[bool, str]:
|
||||
# 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,通过
|
||||
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
|
||||
|
||||
@ -199,6 +228,14 @@ def _check_ip_literal(hostname: str, policy: SSRFPolicy) -> tuple[bool, str]:
|
||||
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 拦截事件(安全审计日志)。
|
||||
|
||||
@ -238,6 +275,9 @@ class SSRFValidator:
|
||||
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 缺少主机名"
|
||||
@ -270,18 +310,38 @@ class SSRFValidator:
|
||||
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}"
|
||||
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 解析会阻塞。"""
|
||||
import asyncio
|
||||
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
@ -289,15 +349,97 @@ class SSRFValidator:
|
||||
|
||||
if loop and loop.is_running():
|
||||
raise RuntimeError("validate_url_sync 不能在已运行的事件循环中调用,请使用 validate_url")
|
||||
return asyncio.run(self.validate_url(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 列表。
|
||||
"""
|
||||
import asyncio
|
||||
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]
|
||||
@ -305,9 +447,7 @@ class SSRFValidator:
|
||||
return cached_ips
|
||||
|
||||
try:
|
||||
addr_infos = await asyncio.to_thread(
|
||||
socket.getaddrinfo, hostname, None, socket.AF_UNSPEC, socket.SOCK_STREAM
|
||||
)
|
||||
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: 调用方应检查空列表
|
||||
|
||||
Loading…
Reference in New Issue
Block a user