ForcePilot/backend/package/yuxi/channel/plugins/dependency_validator.py
Kris dd1857221c feat(channel/plugins): 新增渠道插件系统完整实现
实现了渠道插件的全生命周期管理能力,包括:
1. 插件清单的加载、解析与存储
2. 依赖校验与拓扑排序安装/卸载
3. 插件发现与安全沙箱执行环境
4. 插件注册中心与目录管理能力
2026-05-21 10:27:58 +08:00

288 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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)