288 lines
11 KiB
Python
288 lines
11 KiB
Python
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
import logging
|
|||
|
|
import os
|
|||
|
|
import re
|
|||
|
|
from dataclasses import dataclass, field
|
|||
|
|
from pathlib import Path
|
|||
|
|
from typing import Any
|
|||
|
|
|
|||
|
|
logger = logging.getLogger(__name__)
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass(frozen=True)
|
|||
|
|
class AllowedPackage:
|
|||
|
|
"""白名单中的允许安装包配置。"""
|
|||
|
|
|
|||
|
|
name: str
|
|||
|
|
min_version: str | None = None
|
|||
|
|
max_version: str | None = None
|
|||
|
|
allowed_versions: list[str] | None = None
|
|||
|
|
blocked_versions: list[str] | None = None
|
|||
|
|
|
|||
|
|
|
|||
|
|
_BUILTIN_ALLOWED = {
|
|||
|
|
"aiofiles": AllowedPackage(name="aiofiles", min_version="0.8.0"),
|
|||
|
|
"aiohttp": AllowedPackage(name="aiohttp", min_version="3.8.0"),
|
|||
|
|
"aiosignal": AllowedPackage(name="aiosignal", min_version="1.2.0"),
|
|||
|
|
"asyncpg": AllowedPackage(name="asyncpg", min_version="0.28.0"),
|
|||
|
|
"cachetools": AllowedPackage(name="cachetools", min_version="5.0.0"),
|
|||
|
|
"cryptography": AllowedPackage(name="cryptography", min_version="41.0.0"),
|
|||
|
|
"fastapi": AllowedPackage(name="fastapi", min_version="0.100.0"),
|
|||
|
|
"frozenlist": AllowedPackage(name="frozenlist", min_version="1.3.0"),
|
|||
|
|
"httpx": AllowedPackage(name="httpx", min_version="0.24.0"),
|
|||
|
|
"lxml": AllowedPackage(name="lxml", min_version="4.9.0"),
|
|||
|
|
"multidict": AllowedPackage(name="multidict", min_version="6.0.0"),
|
|||
|
|
"numpy": AllowedPackage(name="numpy", min_version="1.24.0"),
|
|||
|
|
"orjson": AllowedPackage(name="orjson", min_version="3.9.0"),
|
|||
|
|
"pillow": AllowedPackage(name="pillow", min_version="10.0.0"),
|
|||
|
|
"psycopg2-binary": AllowedPackage(name="psycopg2-binary", min_version="2.9.0"),
|
|||
|
|
"pydantic": AllowedPackage(name="pydantic", min_version="2.0.0"),
|
|||
|
|
"pydantic-core": AllowedPackage(name="pydantic-core", min_version="2.0.0"),
|
|||
|
|
"pydantic-settings": AllowedPackage(name="pydantic-settings", min_version="2.0.0"),
|
|||
|
|
"python-dotenv": AllowedPackage(name="python-dotenv", min_version="1.0.0"),
|
|||
|
|
"pyyaml": AllowedPackage(name="pyyaml", min_version="6.0"),
|
|||
|
|
"redis": AllowedPackage(name="redis", min_version="4.5.0"),
|
|||
|
|
"requests": AllowedPackage(name="requests", min_version="2.28.0"),
|
|||
|
|
"sqlalchemy": AllowedPackage(name="sqlalchemy", min_version="2.0.0"),
|
|||
|
|
"starlette": AllowedPackage(name="starlette", min_version="0.27.0"),
|
|||
|
|
"structlog": AllowedPackage(name="structlog", min_version="23.0.0"),
|
|||
|
|
"tenacity": AllowedPackage(name="tenacity", min_version="8.0.0"),
|
|||
|
|
"urllib3": AllowedPackage(name="urllib3", min_version="1.26.0"),
|
|||
|
|
"uvicorn": AllowedPackage(name="uvicorn", min_version="0.22.0"),
|
|||
|
|
"websockets": AllowedPackage(name="websockets", min_version="10.0"),
|
|||
|
|
"yarl": AllowedPackage(name="yarl", min_version="1.8.0"),
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass
|
|||
|
|
class ValidationResult:
|
|||
|
|
"""依赖校验结果。"""
|
|||
|
|
|
|||
|
|
is_valid: bool
|
|||
|
|
package: str
|
|||
|
|
requested_spec: str
|
|||
|
|
errors: list[str] = field(default_factory=list)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class DependencyValidator:
|
|||
|
|
"""渠道插件依赖校验器 —— 防止恶意包安装。
|
|||
|
|
|
|||
|
|
功能:
|
|||
|
|
1. 白名单校验:只允许安装预审批的包
|
|||
|
|
2. 版本约束校验:确保版本符合安全策略
|
|||
|
|
3. 包名规范化:防止依赖混淆攻击
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
_PEP508_NAME_RE = re.compile(r"^[A-Za-z0-9][-A-Za-z0-9._]*$")
|
|||
|
|
_VERSION_SPEC_RE = re.compile(
|
|||
|
|
r"^[A-Za-z0-9][-A-Za-z0-9._]*" # 包名
|
|||
|
|
r"(\s*\[\s*[A-Za-z0-9_,\s-]+\s*\])?" # extras
|
|||
|
|
r"(\s*(==|!=|<=|>=|<|>|~=|===)\s*[^\s;]+)*" # 版本约束
|
|||
|
|
r"(\s*;.*)?$" # 环境标记
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
def __init__(self, allowed_packages: dict[str, AllowedPackage] | None = None):
|
|||
|
|
self._allowed: dict[str, AllowedPackage] = {}
|
|||
|
|
if allowed_packages:
|
|||
|
|
for name, cfg in allowed_packages.items():
|
|||
|
|
self._allowed[self._normalize_name(name)] = cfg
|
|||
|
|
|
|||
|
|
@classmethod
|
|||
|
|
def from_config_file(cls, config_path: str | Path) -> "DependencyValidator":
|
|||
|
|
"""从 JSON 配置文件加载白名单。"""
|
|||
|
|
import json
|
|||
|
|
|
|||
|
|
path = Path(config_path)
|
|||
|
|
if not path.exists():
|
|||
|
|
logger.warning("Dependency whitelist config not found: %s", path)
|
|||
|
|
return cls()
|
|||
|
|
|
|||
|
|
try:
|
|||
|
|
data = json.loads(path.read_text(encoding="utf-8"))
|
|||
|
|
allowed = {}
|
|||
|
|
for item in data.get("allowed_packages", []):
|
|||
|
|
name = item["name"]
|
|||
|
|
allowed[name] = AllowedPackage(
|
|||
|
|
name=name,
|
|||
|
|
min_version=item.get("min_version"),
|
|||
|
|
max_version=item.get("max_version"),
|
|||
|
|
allowed_versions=item.get("allowed_versions"),
|
|||
|
|
blocked_versions=item.get("blocked_versions"),
|
|||
|
|
)
|
|||
|
|
return cls(allowed)
|
|||
|
|
except Exception:
|
|||
|
|
logger.exception("Failed to load dependency whitelist from %s", path)
|
|||
|
|
return cls()
|
|||
|
|
|
|||
|
|
@classmethod
|
|||
|
|
def default_validator(cls) -> "DependencyValidator":
|
|||
|
|
"""获取默认校验器。
|
|||
|
|
|
|||
|
|
加载优先级:
|
|||
|
|
1. 环境变量 YUXI_DEPENDENCY_WHITELIST 指定的文件
|
|||
|
|
2. 统一配置 yuxi.config.config.dependency_whitelist
|
|||
|
|
3. 内置硬编码白名单(降级)
|
|||
|
|
"""
|
|||
|
|
whitelist_file = os.getenv("YUXI_DEPENDENCY_WHITELIST")
|
|||
|
|
if whitelist_file:
|
|||
|
|
path = Path(whitelist_file)
|
|||
|
|
if path.exists():
|
|||
|
|
logger.info("Loading dependency whitelist from env: %s", path)
|
|||
|
|
return cls.from_config_file(path)
|
|||
|
|
|
|||
|
|
try:
|
|||
|
|
from yuxi.config import config
|
|||
|
|
|
|||
|
|
allowed: dict[str, AllowedPackage] = {}
|
|||
|
|
for entry in config.dependency_whitelist:
|
|||
|
|
allowed[entry.name] = AllowedPackage(
|
|||
|
|
name=entry.name,
|
|||
|
|
min_version=entry.min_version,
|
|||
|
|
max_version=entry.max_version,
|
|||
|
|
allowed_versions=entry.allowed_versions,
|
|||
|
|
blocked_versions=entry.blocked_versions,
|
|||
|
|
)
|
|||
|
|
logger.info("Loaded dependency whitelist from unified config (%d packages)", len(allowed))
|
|||
|
|
return cls(allowed)
|
|||
|
|
except Exception:
|
|||
|
|
logger.exception("Failed to load dependency whitelist from unified config, using builtin defaults")
|
|||
|
|
|
|||
|
|
logger.info("Using builtin dependency whitelist defaults (%d packages)", len(_BUILTIN_ALLOWED))
|
|||
|
|
return cls(dict(_BUILTIN_ALLOWED))
|
|||
|
|
|
|||
|
|
def validate(self, dependency_spec: str) -> ValidationResult:
|
|||
|
|
"""校验单个依赖规格是否允许安装。"""
|
|||
|
|
dependency_spec = dependency_spec.strip()
|
|||
|
|
if not dependency_spec:
|
|||
|
|
return ValidationResult(
|
|||
|
|
is_valid=False,
|
|||
|
|
package="",
|
|||
|
|
requested_spec=dependency_spec,
|
|||
|
|
errors=["Empty dependency specification"],
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
if not self._VERSION_SPEC_RE.match(dependency_spec):
|
|||
|
|
return ValidationResult(
|
|||
|
|
is_valid=False,
|
|||
|
|
package="",
|
|||
|
|
requested_spec=dependency_spec,
|
|||
|
|
errors=["Invalid dependency specification format"],
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
package_name = self._extract_package_name(dependency_spec)
|
|||
|
|
normalized = self._normalize_name(package_name)
|
|||
|
|
|
|||
|
|
if normalized not in self._allowed:
|
|||
|
|
return ValidationResult(
|
|||
|
|
is_valid=False,
|
|||
|
|
package=package_name,
|
|||
|
|
requested_spec=dependency_spec,
|
|||
|
|
errors=[f"Package '{package_name}' is not in the allowed whitelist"],
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
allowed = self._allowed[normalized]
|
|||
|
|
version_constraints = self._extract_version_constraints(dependency_spec)
|
|||
|
|
|
|||
|
|
errors = []
|
|||
|
|
for op, ver in version_constraints:
|
|||
|
|
if not self._is_version_allowed(allowed, op, ver):
|
|||
|
|
errors.append(
|
|||
|
|
f"Version constraint '{op}{ver}' violates policy for '{package_name}'"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
if allowed.blocked_versions:
|
|||
|
|
for _, ver in version_constraints:
|
|||
|
|
if ver in allowed.blocked_versions:
|
|||
|
|
errors.append(
|
|||
|
|
f"Version '{ver}' is explicitly blocked for '{package_name}'"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
return ValidationResult(
|
|||
|
|
is_valid=len(errors) == 0,
|
|||
|
|
package=package_name,
|
|||
|
|
requested_spec=dependency_spec,
|
|||
|
|
errors=errors,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
def validate_all(self, dependencies: list[str]) -> list[ValidationResult]:
|
|||
|
|
"""批量校验依赖列表。"""
|
|||
|
|
return [self.validate(dep) for dep in dependencies]
|
|||
|
|
|
|||
|
|
def is_allowed(self, dependency_spec: str) -> bool:
|
|||
|
|
"""快速判断单个依赖是否允许。"""
|
|||
|
|
return self.validate(dependency_spec).is_valid
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _normalize_name(name: str) -> str:
|
|||
|
|
"""规范化包名(PEP 503):转小写,替换 _ 和 . 为 -。"""
|
|||
|
|
return name.lower().replace("_", "-").replace(".", "-")
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _extract_package_name(spec: str) -> str:
|
|||
|
|
"""从依赖规格中提取包名。"""
|
|||
|
|
spec = spec.strip()
|
|||
|
|
match = re.match(r"^([A-Za-z0-9][-A-Za-z0-9._]*)", spec)
|
|||
|
|
if match:
|
|||
|
|
return match.group(1)
|
|||
|
|
return spec.split("[")[0].split(";")[0].strip()
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _extract_version_constraints(spec: str) -> list[tuple[str, str]]:
|
|||
|
|
"""从依赖规格中提取版本约束列表 [(operator, version), ...]。"""
|
|||
|
|
constraints = []
|
|||
|
|
pattern = re.compile(r"(==|!=|<=|>=|<|>|~=|===)\s*([^\s;,]+)")
|
|||
|
|
for match in pattern.finditer(spec):
|
|||
|
|
constraints.append((match.group(1), match.group(2)))
|
|||
|
|
return constraints
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _is_version_allowed(allowed: AllowedPackage, op: str, version: str) -> bool:
|
|||
|
|
"""检查版本约束是否符合白名单策略。"""
|
|||
|
|
if allowed.allowed_versions and version not in allowed.allowed_versions:
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
if allowed.min_version and op in (">=", ">", "=="):
|
|||
|
|
if not DependencyValidator._version_gte(version, allowed.min_version):
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
if allowed.max_version and op in ("<=", "<", "=="):
|
|||
|
|
if not DependencyValidator._version_lte(version, allowed.max_version):
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
return True
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _version_gte(v1: str, v2: str) -> bool:
|
|||
|
|
"""简化版本比较:v1 >= v2。"""
|
|||
|
|
try:
|
|||
|
|
return DependencyValidator._parse_version(v1) >= DependencyValidator._parse_version(v2)
|
|||
|
|
except ValueError:
|
|||
|
|
return True
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _version_lte(v1: str, v2: str) -> bool:
|
|||
|
|
"""简化版本比较:v1 <= v2。"""
|
|||
|
|
try:
|
|||
|
|
return DependencyValidator._parse_version(v1) <= DependencyValidator._parse_version(v2)
|
|||
|
|
except ValueError:
|
|||
|
|
return True
|
|||
|
|
|
|||
|
|
@staticmethod
|
|||
|
|
def _parse_version(version: str) -> tuple[int, ...]:
|
|||
|
|
"""解析版本号为可比较的元组。"""
|
|||
|
|
version = version.strip()
|
|||
|
|
parts = re.split(r"[.-]", version)
|
|||
|
|
result = []
|
|||
|
|
for part in parts:
|
|||
|
|
if part.isdigit():
|
|||
|
|
result.append(int(part))
|
|||
|
|
else:
|
|||
|
|
match = re.match(r"(\d+)", part)
|
|||
|
|
if match:
|
|||
|
|
result.append(int(match.group(1)))
|
|||
|
|
break
|
|||
|
|
if not result:
|
|||
|
|
raise ValueError(f"Invalid version: {version}")
|
|||
|
|
return tuple(result)
|