#!/usr/bin/env python3
"""View-conventions PostToolUse hook.

Reads a Claude Code PostToolUse payload from stdin, identifies the touched
file, and runs the regex checks declared in .claude/conventions/views.json.

Outcome:
  - no violations: exit 0 silently.
  - error-severity matches: exit 0 with JSON {"decision": "block", "reason": ...}
    so Claude is prompted to address the rule on its next turn.
  - advisory-severity matches: exit 0 with JSON
    {"hookSpecificOutput": {"hookEventName": "PostToolUse",
                            "additionalContext": ...}}
    so Claude sees the warning but the tool call proceeds.

Crash safety: any unexpected exception is caught and reported to stderr
with exit 1 (non-blocking), so a hook bug never wedges the agent loop.
"""

from __future__ import annotations

import argparse
import json
import os
import re
import subprocess
import sys
import time
from pathlib import Path

PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
RULE_FILE = PROJECT_ROOT / ".claude" / "conventions" / "views.json"

SKIP_DIRECTIVE_RE = re.compile(
    r"(?P<intro>(?://|\#)\s*lint-skip:(?P<ids1>[a-zA-Z0-9_,\-]+)\s*$)"
    r"|(?P<html><!--\s*lint-skip:(?P<ids2>[a-zA-Z0-9_,\-]+)\s*-->\s*$)"
)

# Alpine/binding/event attributes whose value is a JS *expression*, not display
# text — `x-show`, `:class`, `@click`, `x-transition:enter`, … An echo inside one
# of these is interpolated into JS (e.g. a loop index), so text-only rules such as
# `use-sieroty` must not fire on it. `x-text` / `x-html` are deliberately NOT here:
# they render their value into the DOM as text, so sieroty can still apply.
EXPRESSION_ATTR_RE = re.compile(r'(?:@|:|x-)[\w:.\-]+\s*=\s*"')
EXPRESSION_ATTR_EXCLUDE = {"x-text", "x-html"}


def _in_expression_attr(line: str, start: int, end: int) -> bool:
    """True if the span [start, end) sits inside an Alpine expression-attribute
    value (see EXPRESSION_ATTR_RE). Used to suppress text rules on JS expressions."""
    for m in EXPRESSION_ATTR_RE.finditer(line):
        name = m.group(0).split("=", 1)[0].strip().lstrip("@:")
        if name in EXPRESSION_ATTR_EXCLUDE:
            continue
        val_start = m.end()
        close = line.find('"', val_start)
        if close == -1:
            close = len(line)
        if val_start <= start and end <= close:
            return True
    return False


def glob_to_regex(pattern: str) -> str:
    """Translate a glob with `**` semantics into an anchored regex string.

    `**/` matches zero or more directories. `**` (not followed by `/`)
    matches any sequence including `/`. `*` matches a single path segment.
    """
    out: list[str] = []
    i = 0
    n = len(pattern)
    while i < n:
        c = pattern[i]
        if c == "*":
            if i + 1 < n and pattern[i + 1] == "*":
                if i + 2 < n and pattern[i + 2] == "/":
                    out.append("(?:.*/)?")
                    i += 3
                    continue
                out.append(".*")
                i += 2
                continue
            out.append("[^/]*")
            i += 1
            continue
        if c == "?":
            out.append("[^/]")
        elif c in r".+()[]{}|^$\\":
            out.append("\\" + c)
        else:
            out.append(c)
        i += 1
    return "^" + "".join(out) + "$"


def matches_glob(rel_path: str, pattern: str) -> bool:
    return re.match(glob_to_regex(pattern), rel_path) is not None


def parse_skip_directives(lines: list[str]) -> dict[int, set[str]]:
    """Return {1-based line number: set(rule ids)} for skip directives."""
    out: dict[int, set[str]] = {}
    for idx, line in enumerate(lines, start=1):
        m = SKIP_DIRECTIVE_RE.search(line)
        if not m:
            continue
        raw = m.group("ids1") or m.group("ids2") or ""
        ids = {r.strip() for r in raw.split(",") if r.strip()}
        if ids:
            out[idx] = ids
    return out


def strip_skip_comment(line: str) -> str:
    """Strip the trailing skip-directive so the regex can't match it."""
    return SKIP_DIRECTIVE_RE.sub("", line).rstrip()


def find_violations(rule: dict, lines: list[str], skips: dict[int, set[str]]):
    """Return [(line_no, snippet), ...] for unsuppressed regex matches."""
    rule_id = rule.get("id", "")
    violations: list[tuple[int, str]] = []
    for d in rule.get("detect", []):
        if d.get("type") != "regex":
            continue
        try:
            pattern = re.compile(d["pattern"])
        except (re.error, KeyError, TypeError):
            continue
        contexts = d.get("skipIfContext", []) or []
        skip_expr_attr = bool(d.get("skipInExpressionAttr"))
        for i, line in enumerate(lines, start=1):
            matches = list(pattern.finditer(line))
            if not matches:
                continue
            if rule_id in skips.get(i, set()):
                continue
            if any(token and token in line for token in contexts):
                continue
            # Drop matches that are JS expressions inside Alpine attributes; only
            # flag the line if a genuine (text-context) match remains.
            if skip_expr_attr and not any(
                not _in_expression_attr(line, m.start(), m.end()) for m in matches
            ):
                continue
            violations.append((i, line.strip()))
    return violations


def format_block(rule: dict, violations: list[tuple[int, str]]) -> str:
    head = f"[{rule.get('severity', 'unknown')}] {rule.get('id', '?')} — {rule.get('title', '').strip()}"
    parts = [head.rstrip(" —")]
    for line_no, snippet in violations:
        parts.append(f"  line {line_no}: {snippet}")
    suggest = (rule.get("suggest") or "").strip()
    if suggest:
        parts.append("Suggested fix:")
        parts.append(f"  {suggest}")
    ref = (rule.get("reference") or "").strip()
    if ref:
        parts.append(f"See: {ref}")
    return "\n".join(parts)


def load_payload() -> dict:
    raw = sys.stdin.read()
    if not raw.strip():
        return {}
    try:
        return json.loads(raw)
    except json.JSONDecodeError:
        return {}


def resolve_relative(file_path: str) -> str | None:
    if not file_path:
        return None
    p = Path(file_path)
    if not p.is_absolute():
        p = (PROJECT_ROOT / p).resolve()
    try:
        rel = p.relative_to(PROJECT_ROOT)
    except ValueError:
        return None
    return str(rel).replace(os.sep, "/")


def load_rules() -> list[dict]:
    if not RULE_FILE.exists():
        return []
    try:
        doc = json.loads(RULE_FILE.read_text(encoding="utf-8"))
    except (OSError, json.JSONDecodeError):
        return []
    rules = doc.get("rules") or []
    return rules if isinstance(rules, list) else []


def check_content(rel_path: str, content: str, rules: list[dict]):
    """Run all matching rules against in-memory content.

    Returns (errors, advisories, skip_records). Shared by the PostToolUse hook
    (which reads the file from disk via check_file) and the CLI (which may read
    the staged blob content from the git index).
    """
    matching = [r for r in rules if any(matches_glob(rel_path, p) for p in (r.get("files") or []))]
    if not matching:
        return [], [], []

    raw_lines = content.splitlines()
    skips = parse_skip_directives(raw_lines)
    cleaned = [strip_skip_comment(l) for l in raw_lines]

    errors: list[str] = []
    advisories: list[str] = []
    for rule in matching:
        violations = find_violations(rule, cleaned, skips)
        if not violations:
            continue
        block = format_block(rule, violations)
        if rule.get("severity") == "error":
            errors.append(block)
        else:
            advisories.append(block)

    skip_records: list[str] = []
    matching_ids = {r.get("id") for r in matching}
    for line_no in sorted(skips):
        for rid in sorted(skips[line_no]):
            if rid in matching_ids:
                skip_records.append(f"[skip] line {line_no}: rule {rid} suppressed by directive")

    return errors, advisories, skip_records


def check_file(rel_path: str, abs_path: Path, rules: list[dict]):
    """Run all matching rules against a file on disk. Returns (errors, advisories, skip_records)."""
    try:
        content = abs_path.read_text(encoding="utf-8")
    except (OSError, UnicodeDecodeError):
        return [], [], []
    return check_content(rel_path, content, rules)


# ---------------------------------------------------------------------------
# CLI mode — commit-time / CI enforcement (reuses the same rule engine)
# ---------------------------------------------------------------------------

def _git_output(args: list[str]) -> str | None:
    """Run a git command at PROJECT_ROOT; return stdout, or None on failure."""
    try:
        proc = subprocess.run(
            ["git", *args],
            cwd=str(PROJECT_ROOT),
            capture_output=True,
            text=True,
            timeout=15,
        )
    except (OSError, subprocess.SubprocessError):
        return None
    if proc.returncode != 0:
        return None
    return proc.stdout


def staged_php_files() -> list[str]:
    """Repo-relative paths of staged (added/copied/modified) *.php files."""
    out = _git_output(["diff", "--cached", "--name-only", "--diff-filter=ACM", "-z"])
    if out is None:
        return []
    return [p for p in out.split("\0") if p and p.endswith(".php")]


def staged_content(rel_path: str) -> str | None:
    """Content of the staged (index) version of rel_path, or None on failure."""
    return _git_output(["show", f":{rel_path}"])


def collect_targets(args) -> list[tuple[str, str]]:
    """Return [(rel_path, content), ...] for the requested inputs.

    `--staged` reads the index (so partially-staged files are checked as they
    will be committed); `--files` reads the working tree.
    """
    targets: list[tuple[str, str]] = []
    if args.staged:
        for rel in staged_php_files():
            content = staged_content(rel)
            if content is not None:
                targets.append((rel, content))
    for f in args.files:
        rel = resolve_relative(f)
        if rel is None:
            continue
        abs_path = (PROJECT_ROOT / rel).resolve()
        if not abs_path.is_file():
            continue
        try:
            content = abs_path.read_text(encoding="utf-8")
        except (OSError, UnicodeDecodeError):
            continue
        targets.append((rel, content))
    return targets


def run_cli(argv: list[str]) -> int:
    """Check staged and/or explicit files; return 1 only on error-severity violations.

    Advisory findings are printed but never fail the run. Honors the same
    `lint-skip:` directives as the PostToolUse hook.
    """
    parser = argparse.ArgumentParser(
        prog="view-conventions-check",
        description="Enforce view conventions on staged or explicit files.",
    )
    parser.add_argument(
        "--staged", action="store_true",
        help="check the staged (index) version of changed *.php files",
    )
    parser.add_argument(
        "--files", nargs="*", default=[], metavar="PATH",
        help="explicit file paths to check (read from the working tree)",
    )
    args = parser.parse_args(argv)

    rules = load_rules()
    if not rules:
        return 0

    error_count = 0
    report: list[str] = []
    for rel, content in collect_targets(args):
        errors, advisories, skip_records = check_content(rel, content, rules)
        if not errors and not advisories:
            continue
        report.append(f"View-conventions check on {rel}")
        if errors:
            report.append("  Blocking violations:")
            report.extend(errors)
            error_count += len(errors)
        if advisories:
            report.append("  Advisory warnings:")
            report.extend(advisories)
        if skip_records:
            report.extend(skip_records)
        report.append("")

    if report:
        sys.stderr.write("\n".join(report) + "\n")
    if error_count:
        sys.stderr.write(
            f"\n✖ {error_count} blocking view-convention rule(s) violated. "
            "Fix them, add a `lint-skip:<rule>` directive on the offending line, "
            "or bypass with `git commit --no-verify`.\n"
        )
        return 1
    return 0


def main() -> int:
    payload = load_payload()
    tool_name = payload.get("tool_name", "")
    if tool_name not in ("Edit", "Write", "MultiEdit"):
        return 0

    tool_input = payload.get("tool_input") or {}
    rel = resolve_relative(tool_input.get("file_path", ""))
    if rel is None:
        return 0

    rules = load_rules()
    if not rules:
        return 0

    abs_path = (PROJECT_ROOT / rel).resolve()
    if not abs_path.is_file():
        return 0

    errors, advisories, skip_records = check_file(rel, abs_path, rules)

    if not errors and not advisories:
        return 0

    sections: list[str] = [f"View-conventions check on {rel}"]
    if errors:
        sections.append("")
        sections.append("Blocking violations — must fix before next edit:")
        sections.extend(errors)
    if advisories:
        sections.append("")
        sections.append("Advisory warnings:")
        sections.extend(advisories)
    if skip_records:
        sections.append("")
        sections.extend(skip_records)
    message = "\n".join(sections)

    if errors:
        out = {"decision": "block", "reason": message}
    else:
        out = {
            "hookSpecificOutput": {
                "hookEventName": "PostToolUse",
                "additionalContext": message,
            }
        }
    sys.stdout.write(json.dumps(out))
    return 0


if __name__ == "__main__":
    started = time.perf_counter()
    argv = sys.argv[1:]
    cli_mode = bool(argv)  # any arguments → CLI; no args → stdin PostToolUse hook
    try:
        code = run_cli(argv) if cli_mode else main()
    except Exception as exc:  # never wedge the caller
        sys.stderr.write(f"view-conventions-check: internal error: {exc}\n")
        # CLI: exit 2 so the git pre-commit hook can fail open (allow the commit
        # with a warning) on a checker bug. Hook mode: exit 1 (non-blocking for
        # Claude), as before.
        sys.exit(2 if cli_mode else 1)
    elapsed_ms = (time.perf_counter() - started) * 1000
    if os.environ.get("VIEW_CONVENTIONS_DEBUG"):
        sys.stderr.write(f"view-conventions-check: {elapsed_ms:.1f} ms\n")
    sys.exit(code)
