from __future__ import annotations import logging import re from abc import ABC, abstractmethod from dataclasses import dataclass from fnmatch import fnmatch from pathlib import Path import semantic_version from semantic_version.base import AllOf, AnyOf, Range, Version from vuln_scan.core.models import ( Dependency, Ecosystem, EcosystemParseResult, SecurityVulnerability, ) logger = logging.getLogger(__name__) class EcosystemHandler(ABC): """ Base class for all ecosystem handlers. Responsibilities: - Identify supported manifests - Parse manifests into dependencies - Determine vulnerability matching logic """ id: str ecosystem: Ecosystem # ----------------------------- # Interval Engine (SHARED) # ----------------------------- @dataclass class _Interval: lower: semantic_version.Version | None upper: semantic_version.Version | None include_lower: bool include_upper: bool def overlaps(self, other: EcosystemHandler._Interval) -> bool: # Resolve the effective lower bound and its inclusivity if self.lower is None: lower = other.lower inc_lower = other.include_lower elif other.lower is None or self.lower > other.lower: lower = self.lower inc_lower = self.include_lower elif other.lower > self.lower: lower = other.lower inc_lower = other.include_lower else: # equal lowers — both must include it lower = self.lower inc_lower = self.include_lower and other.include_lower # Resolve the effective upper bound and its inclusivity if self.upper is None: upper = other.upper inc_upper = other.include_upper elif other.upper is None or self.upper < other.upper: upper = self.upper inc_upper = self.include_upper elif other.upper < self.upper: upper = other.upper inc_upper = other.include_upper else: # equal uppers — both must include it upper = self.upper inc_upper = self.include_upper and other.include_upper if lower is None or upper is None: return True if lower < upper: return True if lower == upper: return inc_lower and inc_upper return False def _range_to_interval(self, r: Range) -> _Interval: op, v = r.operator, r.target if op == ">": return self._Interval(v, None, False, False) if op == ">=": return self._Interval(v, None, True, False) if op == "<": return self._Interval(None, v, False, False) if op == "<=": return self._Interval(None, v, False, True) if op == "==": return self._Interval(v, v, True, True) raise ValueError(f"Unsupported operator: {op}") def _intersect_intervals(self, intervals: list[_Interval]) -> _Interval: result = intervals[0] for other in intervals[1:]: # Resolve lower bound — keep the tighter (higher) lower if result.lower is None: lower = other.lower inc_lower = other.include_lower elif other.lower is None or result.lower > other.lower: lower = result.lower inc_lower = result.include_lower elif other.lower > result.lower: lower = other.lower inc_lower = other.include_lower else: # equal — both must include lower = result.lower inc_lower = result.include_lower and other.include_lower # Resolve upper bound — keep the tighter (lower) upper if result.upper is None: upper = other.upper inc_upper = other.include_upper elif other.upper is None or result.upper < other.upper: upper = result.upper inc_upper = result.include_upper elif other.upper < result.upper: upper = other.upper inc_upper = other.include_upper else: # equal — both must include upper = result.upper inc_upper = result.include_upper and other.include_upper result = self._Interval(lower, upper, inc_lower, inc_upper) return result def _clause_to_intervals(self, clause: Range | AllOf | AnyOf) -> list[_Interval]: if isinstance(clause, Range): return [self._range_to_interval(clause)] if isinstance(clause, AllOf): intervals = [] for c in clause.clauses: intervals.extend(self._clause_to_intervals(c)) return [self._intersect_intervals(intervals)] if isinstance(clause, AnyOf): result = [] for c in clause.clauses: result.extend(self._clause_to_intervals(c)) return result return [] def _ranges_overlap(self, r1: str, r2: str) -> bool: try: s1 = self._build_spec(r1) s2 = self._build_spec(r2) for i1 in self._clause_to_intervals(s1.clause): for i2 in self._clause_to_intervals(s2.clause): if i1.overlaps(i2): return True return False except Exception: logger.warning("Range overlap failed", exc_info=True) return True def _version_in_range(self, version: Version, vuln_range: str) -> bool: try: return bool(self._build_spec(vuln_range).match(version)) except Exception: logger.warning("Version match failed", exc_info=True) return True # ----------------------------- # Abstract hooks # ----------------------------- def _build_spec(self, range_str: str) -> semantic_version.NpmSpec: raise NotImplementedError def _normalize_version(self, version: str) -> str: return version def _normalize_vuln_range(self, spec: str) -> str: spec = (spec or "").strip() if not spec: return "" if re.fullmatch(r"\d+(\.\d+)*", spec): return f"=={spec}" if spec.startswith("=") and not spec.startswith("=="): return f"=={spec.lstrip('=').strip()}" return spec # ----------------------------- # Manifest logic (unchanged) # ----------------------------- @property def lockfile_names(self) -> set[str]: """ Exact lockfile names. """ return set() @property def manifest_names(self) -> set[str]: """ Exact manifest filenames. """ return set() @property def manifest_globs(self) -> set[str]: """ Glob patterns for manifests. Example: {"*requirements*.txt"} """ return set() @property def all_target_filenames(self) -> set[str]: return self.lockfile_names | self.manifest_names | self.manifest_globs # ------------------------------------------------------- # Manifest detection # ------------------------------------------------------- def supports_manifest(self, manifest_path: str) -> bool: filename = Path(manifest_path).name if filename in self.lockfile_names: return True if filename in self.manifest_names: return True return any(fnmatch(filename, pattern) for pattern in self.manifest_globs) def is_lockfile(self, filename: str) -> bool: return filename in self.lockfile_names def is_manifest(self, filename: str) -> bool: if filename in self.manifest_names: return True return any(fnmatch(filename, pattern) for pattern in self.manifest_globs) # ------------------------------------------------------- # Parsing # ------------------------------------------------------- @abstractmethod def parse_manifest( self, manifest_path: str, full_path: Path, ) -> EcosystemParseResult: """ Parse a manifest or lockfile from disk. """ raise NotImplementedError @abstractmethod def parse_manifest_content( self, manifest_path: str, content: str, ) -> EcosystemParseResult: """ Parse manifest content (used for GitHub API or memory). """ raise NotImplementedError # ------------------------------------------------------- # Vulnerability matching # ------------------------------------------------------- @abstractmethod def is_vulnerable( self, dependency: Dependency, alert: SecurityVulnerability, ) -> tuple[bool, str]: """ Returns: (is_vulnerable, confidence) confidence: high → exact version match low → spec/range fallback """ raise NotImplementedError