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)