#!/usr/bin/env python3
"""Prototype pollution pre-production check script.

Scans JavaScript source files for common patterns that may indicate
prototype pollution risk, including:
- dangerous property names (__proto__, constructor, prototype)
- object merges and extensions
- dynamic property writes
- recursive deep copy / merge helpers

This script is intended as a lightweight review aid, not a full security
scanner. Review all findings manually before taking action.
"""

from __future__ import annotations

import argparse
import os
import re
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable, List, Sequence


DANGEROUS_KEYS = ("__proto__", "prototype", "constructor")
DEFAULT_EXTENSIONS = {".js", ".mjs", ".cjs", ".jsx", ".ts", ".tsx"}


@dataclass
class Finding:
    path: str
    line: int
    category: str
    message: str
    excerpt: str


PATTERNS = [
    ("dangerous-key", re.compile(r'(?i)(__proto__|\bconstructor\b|\bprototype\b)')),
    ("object-assign", re.compile(r'\bObject\.assign\s*\(')),
    ("spread-merge", re.compile(r'\{[^\n]*\.\.\.[^\n]*\}')),
    ("deep-merge", re.compile(r'\b(deepMerge|mergeDeep|recursiveMerge|extendDeep)\b', re.I)),
    ("dynamic-write", re.compile(r'\[[^\]]+\]\s*=')),
    ("path-split-write", re.compile(r'\b(path|key|prop|property)\s*\.\s*split\s*\(')),
    ("object-create", re.compile(r'\bObject\.create\s*\(\s*null\s*\)')),
]


def iter_files(paths: Sequence[str], recursive: bool, exts: set[str]) -> Iterable[Path]:
    for raw in paths:
        p = Path(raw)
        if not p.exists():
            continue
        if p.is_file():
            if p.suffix.lower() in exts:
                yield p
            continue
        if p.is_dir() and recursive:
            for child in p.rglob("*"):
                if child.is_file() and child.suffix.lower() in exts:
                    yield child
        elif p.is_dir():
            for child in p.iterdir():
                if child.is_file() and child.suffix.lower() in exts:
                    yield child


def scan_file(path: Path) -> List[Finding]:
    findings: List[Finding] = []
    try:
        text = path.read_text(encoding="utf-8", errors="replace")
    except OSError as exc:
        findings.append(
            Finding(str(path), 0, "read-error", f"Could not read file: {exc}", "")
        )
        return findings

    lines = text.splitlines()
    for idx, line in enumerate(lines, start=1):
        for category, pattern in PATTERNS:
            if pattern.search(line):
                msg = {
                    "dangerous-key": "Potential prototype-related property name found.",
                    "object-assign": "Object.assign() found; verify input is trusted and keys are filtered.",
                    "spread-merge": "Object spread detected; review whether untrusted nested data is being merged.",
                    "deep-merge": "Deep merge helper reference found; inspect for key filtering and prototype safeguards.",
                    "dynamic-write": "Dynamic property write detected; verify path validation and denylist checks.",
                    "path-split-write": "Path-splitting logic found; ensure prototype-related keys are rejected.",
                    "object-create": "Null-prototype object found; this is often safer for dictionaries.",
                }.get(category, "Potentially relevant pattern found.")
                findings.append(
                    Finding(str(path), idx, category, msg, line.strip())
                )
    return findings


def print_findings(findings: Sequence[Finding]) -> None:
    if not findings:
        print("No matching patterns found.")
        return
    for f in findings:
        print(f"{f.path}:{f.line}: [{f.category}] {f.message}")
        if f.excerpt:
            print(f"  {f.excerpt}")


def main() -> int:
    parser = argparse.ArgumentParser(
        description="Scan JavaScript source for prototype pollution risk patterns."
    )
    parser.add_argument(
        "paths",
        nargs="+",
        help="One or more files or directories to scan.",
    )
    parser.add_argument(
        "--recursive",
        action="store_true",
        help="Recurse into directories.",
    )
    parser.add_argument(
        "--include",
        nargs="*",
        default=sorted(DEFAULT_EXTENSIONS),
        help="File extensions to include, e.g. .js .ts .jsx.",
    )
    parser.add_argument(
        "--fail-on-findings",
        action="store_true",
        help="Exit with code 1 if any findings are reported.",
    )
    args = parser.parse_args()

    exts = {e if e.startswith(".") else f".{e}" for e in args.include}
    findings: List[Finding] = []

    for file_path in iter_files(args.paths, args.recursive, exts):
        findings.extend(scan_file(file_path))

    print_findings(findings)

    if args.fail_on_findings and findings:
        return 1
    return 0


if __name__ == "__main__":
    raise SystemExit(main())