ForcePilot/backend/package/yuxi/channel/plugins/dependency_validator.py

288 lines
11 KiB
Python
Raw Normal View History

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)