#!/usr/bin/env python3
"""Pre-tool hook for Claude Code, Codex, and Gemini CLI.

Reads the tool-call JSON on stdin, asks aiproxy-supervisor, and denies the
call when the supervisor denies it. Tools with no file, command, or URL are
left alone.
"""
from __future__ import annotations

import json
import os
import re
import shlex
import subprocess
import sys
from pathlib import Path
from urllib.parse import urlparse

WRITE_TOOLS = {
    "write",
    "edit",
    "multiedit",
    "notebookedit",
    "write_file",
    "replace",
    "apply_patch",
    "search_replace",
}
PATH_KEYS = ("file_path", "path", "target_file", "notebook_path", "file")


def home() -> Path:
    return Path(os.environ.get("HOME", str(Path.home())))


def supervisor() -> Path | None:
    env = os.environ.get("AIPROXY_SUPERVISOR")
    candidates = []
    if env:
        candidates.append(Path(env))
    candidates.extend(
        [
            home() / ".local/aiproxy/bin/aiproxy-supervisor",
            Path(__file__).resolve().parent / "aiproxy-supervisor",
        ]
    )
    for c in candidates:
        for path in (c, c.with_suffix(".exe") if c.suffix == "" else c):
            if path.is_file() and (os.access(path, os.X_OK) or os.name == "nt"):
                return path
    from shutil import which

    found = which("aiproxy-supervisor")
    return Path(found) if found else None


def policy_and_keys(root_hint: Path | None) -> tuple[Path, Path]:
    policy = os.environ.get("AIPROXY_POLICY")
    keys = os.environ.get("AIPROXY_KEYS")
    roots = []
    if root_hint:
        roots.append(root_hint)
    roots.append(home() / ".local/aiproxy")
    roots.append(Path.cwd())
    pol = Path(policy) if policy else None
    key = Path(keys) if keys else None
    if pol is None:
        for r in roots:
            cand = r / "policy/max-boundary.yaml"
            if cand.is_file():
                pol = cand
                break
    if pol is None:
        pol = home() / ".local/aiproxy/policy/max-boundary.yaml"
    if key is None:
        for r in roots:
            cand = r / "var/keys"
            if cand.is_dir():
                key = cand
                break
        if key is None:
            key = home() / ".local/aiproxy/var/keys"
    return pol, key


def ensure_keys(binary: Path, keys: Path) -> None:
    if (keys / "ed25519.sk").is_file() or (keys / "mode").is_file():
        return
    keys.mkdir(parents=True, exist_ok=True)
    subprocess.run([str(binary), "--keys", str(keys), "init-keys"], check=False, capture_output=True, text=True)


def run_check(binary: Path, policy: Path, keys: Path, args: list[str]) -> tuple[str, str]:
    proc = subprocess.run(
        [str(binary), "--policy", str(policy), "--keys", str(keys), *args],
        capture_output=True,
        text=True,
        timeout=20,
    )
    raw = (proc.stdout or "") + (proc.stderr or "")
    i = raw.find("{")
    if i < 0:
        return "deny", "supervisor returned no decision"
    try:
        doc = json.loads(raw[i:])
    except json.JSONDecodeError:
        return "deny", "supervisor returned unreadable output"
    decision = doc.get("decision") or "deny"
    reason = doc.get("reason") or doc.get("detail") or decision
    return decision, str(reason)


def first_binary(command: str) -> str:
    try:
        parts = shlex.split(command)
    except ValueError:
        parts = command.split()
    if not parts:
        return ""
    if parts[0] in {"sudo", "command", "exec", "time"} and len(parts) > 1:
        return parts[1]
    return parts[0]


def paths_in(command: str) -> list[str]:
    found = re.findall(r"(?:(?:~|/|\.\.?/)[^\s'\"|;<>]+)", command)
    out = []
    for p in found:
        if p.startswith("~"):
            p = str(home() / p[1:].lstrip("/"))
        out.append(p)
    return out


def checks_for(event: dict) -> list[tuple[str, list[str]]]:
    tool = str(event.get("tool_name") or event.get("toolName") or "")
    inp = event.get("tool_input") or event.get("toolInput") or {}
    if not isinstance(inp, dict):
        inp = {}
    write = tool.lower() in WRITE_TOOLS or any(w in tool.lower() for w in ("write", "edit", "patch", "replace"))
    checks: list[tuple[str, list[str]]] = []
    for key in PATH_KEYS:
        val = inp.get(key)
        if isinstance(val, str) and val.strip():
            checks.append(("fs", ["check-fs", val, *(["--write"] if write else [])]))
    command = inp.get("command") or inp.get("cmd") or ""
    if isinstance(command, str) and command.strip():
        binary = first_binary(command)
        if binary:
            checks.append(("exec", ["check-exec", binary]))
        for path in paths_in(command):
            destructive = bool(re.search(r"\b(rm|mv|cp|tee|dd|truncate)\b", command)) or ">" in command
            checks.append(("fs", ["check-fs", path, *(["--write"] if destructive else [])]))
        urls = re.findall(r"https?://[^\s'\"|;]+", command)
    else:
        urls = []
    url = inp.get("url") or inp.get("uri") or ""
    if isinstance(url, str) and url:
        urls.append(url)
    for url in urls:
        parsed = urlparse(url)
        if not parsed.hostname:
            continue
        port = parsed.port or (443 if parsed.scheme == "https" else 80)
        method = str(inp.get("method") or "GET")
        body = command if isinstance(command, str) else ""
        checks.append(
            (
                "net",
                [
                    "check-net",
                    parsed.hostname,
                    method,
                    "--port",
                    str(port),
                    "--path",
                    parsed.path or "/",
                    "--body",
                    body[:4000],
                ],
            )
        )
    return checks


def emit(decision: str, reason: str) -> int:
    if decision != "deny":
        return 0
    payload = {
        "decision": "deny",
        "reason": reason,
        "hookSpecificOutput": {
            "hookEventName": "PreToolUse",
            "permissionDecision": "deny",
            "permissionDecisionReason": reason,
        },
    }
    sys.stdout.write(json.dumps(payload) + "\n")
    sys.stderr.write(reason + "\n")
    return 2


def main() -> int:
    raw = sys.stdin.read()
    try:
        event = json.loads(raw) if raw.strip() else {}
    except json.JSONDecodeError:
        event = {}
    binary = supervisor()
    if binary is None:
        return emit("deny", "ai-proxy supervisor is not installed on this computer")
    policy, keys = policy_and_keys(binary.parent.parent if binary.parent.name == "bin" else None)
    if not policy.is_file():
        return emit("deny", f"ai-proxy policy is missing at {policy}")
    ensure_keys(binary, keys)
    for _kind, args in checks_for(event):
        try:
            decision, reason = run_check(binary, policy, keys, args)
        except Exception as exc:
            return emit("deny", f"ai-proxy supervisor failed: {exc}")
        if decision == "deny":
            return emit("deny", f"ai-proxy denied this action ({reason})")
    return 0


if __name__ == "__main__":
    try:
        raise SystemExit(main())
    except BrokenPipeError:
        raise SystemExit(0)
