#!/usr/bin/env python3
"""gh-cached: TTL-based LRU cache wrapper for the GitHub CLI.

Drop-in replacement for `gh` that caches read-only command responses
to reduce API calls. Mutations bypass the cache and invalidate
related entries.

Cache is file-backed in /tmp/gh-cache/ for cross-process sharing, scoped into
a per-repo subdirectory (#5224) so two different repos on the same host never
share cache entries even when they issue textually identical `gh` commands.

Usage:
    gh-cached issue view 42 --json labels
    gh-cached pr list --label "loom:review-requested" --state open
    gh-cached --no-cache issue view 42 --json labels  # bypass cache
    gh-cached --clear-cache                            # clear all cached entries
    gh-cached --cache-stats                            # show hit/miss statistics

Environment:
    GH_CACHE_DIR       Cache directory base (default: /tmp/gh-cache); the
                        actual cache lives under a per-repo subdirectory of
                        this path — see resolve_repo_id() below.
    GH_CACHE_TTL       Default TTL in seconds (default: 30)
    GH_CACHE_MAX_SIZE  Max cached entries (default: 256)
    GH_CACHE_DISABLE   Set to "1" to disable caching entirely
    GH_CACHE_DEBUG     Set to "1" for debug logging to stderr
    LOOM_GH_BIN        The `gh` binary to execute (default: `gh` from PATH),
                        for parity with loom-daemon's own LOOM_GH_BIN.
    GH_CACHE_OUTCOME_LOG  Optional path; each invocation appends one JSON
                        record carrying `x-loom-cache: hit|miss|revalidated|
                        bypass` so downstream views can tell client-cache
                        hits from real gateway requests. The last outcome is
                        also kept as `last_outcome` in `_stats.json`.
    GH_CACHE_REPO_ID   Override the resolved repo identity used to scope
                        GH_CACHE_DIR (default: `git rev-parse
                        --show-toplevel`). Mainly for tests and non-git
                        working directories that still want isolation.
"""

from __future__ import annotations

import hashlib
import json
import os
import subprocess
import sys
import time

# ─── Configuration ──────────────────────────────────────────────────────────


def resolve_repo_id() -> str:
    """Resolve a stable per-host-repo identity used to scope CACHE_DIR (#5224).

    Two different repos on the same host issuing the textually identical `gh`
    command must never share a cache entry — `cache_key()` alone only hashes
    the `gh` argv, which says nothing about which repo it targets (most reads
    routed through this wrapper have no explicit `--repo` and resolve via the
    cwd's git remote). Scoping CACHE_DIR itself (rather than only extending
    cache_key()'s hash input) also naturally repo-isolates enforce_max_size(),
    invalidate_for_resource(), invalidate_all_for_type(), and clear_cache(),
    since every one of them already operates on CACHE_DIR as-is.

    Uses `git rev-parse --show-toplevel` — cheap, local, no network call —
    mirroring the existing pattern of resolving $GH_READ once per session via
    a local probe (defaults/docs/gh-cached.md). GH_CACHE_REPO_ID overrides the
    resolved identity directly (used by tests, and as an escape hatch for
    non-git working directories that still want isolation).
    """
    override = os.environ.get("GH_CACHE_REPO_ID", "")
    if override:
        return override
    try:
        proc = subprocess.run(
            ["git", "rev-parse", "--show-toplevel"],
            capture_output=True,
            text=True,
            timeout=5,
        )
        if proc.returncode == 0:
            top = proc.stdout.strip()
            if top:
                return top
    except (OSError, subprocess.TimeoutExpired):
        pass
    # Not inside a git repo (or git unavailable): fall back to a fixed
    # sentinel so behavior is deterministic rather than silently unscoped.
    return "_no-repo"


CACHE_DIR_BASE = os.environ.get("GH_CACHE_DIR", "/tmp/gh-cache")
REPO_ID = resolve_repo_id()
CACHE_DIR = os.path.join(CACHE_DIR_BASE, hashlib.sha256(REPO_ID.encode()).hexdigest()[:16])
DEFAULT_TTL = int(os.environ.get("GH_CACHE_TTL", "30"))
MAX_CACHE_SIZE = int(os.environ.get("GH_CACHE_MAX_SIZE", "256"))
CACHE_DISABLED = os.environ.get("GH_CACHE_DISABLE", "") == "1"
DEBUG = os.environ.get("GH_CACHE_DEBUG", "") == "1"

# TTL overrides by command type (seconds).
#
# These are the hot polling shapes the sweep/judge/champion skills route through
# this wrapper (#4667). They all track DEFAULT_TTL so that GH_CACHE_TTL is a real
# knob: previously these were hard-coded 30s, which silently made GH_CACHE_TTL a
# no-op for exactly the commands that matter (it only applied to the shapes that
# fell through to the default). Unset GH_CACHE_TTL still yields 30s, so this is
# behavior-identical unless an operator deliberately tunes the window.
TTL_BY_COMMAND = {
    ("issue", "view"):   DEFAULT_TTL,
    ("issue", "list"):   DEFAULT_TTL,
    ("pr", "view"):      DEFAULT_TTL,
    ("pr", "list"):      DEFAULT_TTL,
    ("api",):            DEFAULT_TTL,
}

# Commands that are read-only and safe to cache
CACHEABLE_SUBCOMMANDS = frozenset({"view", "list", "search", "status"})

# Commands that mutate state — bypass cache and invalidate
MUTATION_SUBCOMMANDS = frozenset({
    "edit", "create", "delete", "close", "reopen",
    "merge", "review", "comment", "label",
})

# Top-level gh commands that are never cached
PASSTHROUGH_COMMANDS = frozenset({
    "auth", "config", "ssh-key", "gpg-key", "secret",
    "repo", "gist", "extension", "alias", "completion",
    "help", "--help", "-h", "--version",
})

# ─── Cache Implementation ───────────────────────────────────────────────────

def debug(msg: str) -> None:
    if DEBUG:
        print(f"[gh-cached] {msg}", file=sys.stderr)


def ensure_cache_dir() -> None:
    os.makedirs(CACHE_DIR, mode=0o700, exist_ok=True)


def cache_key(args: list[str]) -> str:
    """Generate a cache key from the gh command arguments."""
    raw = " ".join(args)
    return hashlib.sha256(raw.encode()).hexdigest()[:16]


def cache_path(key: str) -> str:
    return os.path.join(CACHE_DIR, f"{key}.json")


def stats_path() -> str:
    return os.path.join(CACHE_DIR, "_stats.json")


def read_stats() -> dict:
    path = stats_path()
    try:
        with open(path) as f:
            return json.load(f)
    except (FileNotFoundError, json.JSONDecodeError):
        return {"hits": 0, "misses": 0, "bypasses": 0, "invalidations": 0, "etag": 0}


def write_stats(stats: dict) -> None:
    ensure_cache_dir()
    path = stats_path()
    try:
        with open(path, "w") as f:
            json.dump(stats, f)
    except OSError:
        pass


def increment_stat(stat_name: str) -> None:
    stats = read_stats()
    stats[stat_name] = stats.get(stat_name, 0) + 1
    write_stats(stats)


def record_outcome(outcome: str, args: list[str]) -> None:
    """Record the `x-loom-cache` outcome (hit|miss|revalidated|bypass).

    A client-side cache in front of the egress gateway is fine; invisible
    stacking is not (loom#9953) — this keeps cache hits distinguishable.
    """
    stats = read_stats()
    stats["last_outcome"] = outcome
    write_stats(stats)
    log_path = os.environ.get("GH_CACHE_OUTCOME_LOG")
    if log_path:
        try:
            with open(log_path, "a") as f:
                f.write(json.dumps({"x-loom-cache": outcome, "args": args}) + "\n")
        except OSError:
            pass


def cache_get(key: str) -> tuple[str | None, int | None]:
    """Read a cached entry. Returns (stdout, returncode) or (None, None) if miss."""
    path = cache_path(key)
    try:
        with open(path) as f:
            entry = json.load(f)
        if time.time() - entry["time"] > entry["ttl"]:
            debug(f"EXPIRED key={key}")
            os.unlink(path)
            return None, None
        # Update access time for LRU
        entry["accessed"] = time.time()
        with open(path, "w") as f:
            json.dump(entry, f)
        return entry["stdout"], entry["returncode"]
    except (FileNotFoundError, json.JSONDecodeError, KeyError):
        return None, None


def cache_put(key: str, stdout: str, returncode: int, ttl: int, args: list[str] | None = None) -> None:
    """Write a cache entry."""
    ensure_cache_dir()
    enforce_max_size()
    entry = {
        "time": time.time(),
        "accessed": time.time(),
        "ttl": ttl,
        "stdout": stdout,
        "returncode": returncode,
        "args": args or [],
    }
    path = cache_path(key)
    try:
        with open(path, "w") as f:
            json.dump(entry, f)
    except OSError:
        pass


def enforce_max_size() -> None:
    """Evict least-recently-accessed entries if cache exceeds max size."""
    try:
        entries = []
        for name in os.listdir(CACHE_DIR):
            if name.startswith("_") or not name.endswith(".json"):
                continue
            path = os.path.join(CACHE_DIR, name)
            try:
                with open(path) as f:
                    data = json.load(f)
                entries.append((data.get("accessed", 0), path))
            except (json.JSONDecodeError, OSError):
                # Corrupted entry — remove it
                os.unlink(path)

        if len(entries) <= MAX_CACHE_SIZE:
            return

        # Sort by access time, evict oldest
        entries.sort(key=lambda x: x[0])
        evict_count = len(entries) - MAX_CACHE_SIZE
        for _, path in entries[:evict_count]:
            debug(f"EVICT {os.path.basename(path)}")
            os.unlink(path)
    except OSError:
        pass


def invalidate_for_resource(resource_type: str, resource_id: str | None) -> None:
    """Invalidate cached entries related to a resource.

    After a mutation like `gh issue edit 42 ...`, we invalidate any cached
    entries whose original command args reference the same resource type and id.
    Falls back to checking stdout for backward compatibility with cache entries
    that predate the args field.
    """
    if not resource_id:
        return

    debug(f"INVALIDATE {resource_type} {resource_id}")
    count = 0
    try:
        for name in os.listdir(CACHE_DIR):
            if name.startswith("_") or not name.endswith(".json"):
                continue
            path = os.path.join(CACHE_DIR, name)
            try:
                with open(path) as f:
                    data = json.load(f)
                args = data.get("args", [])
                # Primary: check if the cached command's args reference this resource
                if args and resource_type in args and resource_id in args:
                    debug(f"INVALIDATE (args match) {name}")
                    os.unlink(path)
                    count += 1
                    continue
                # Fallback: check stdout for backward compatibility with
                # cache entries that don't have stored args
                stdout = data.get("stdout", "")
                if resource_id in stdout:
                    debug(f"INVALIDATE (stdout match) {name}")
                    os.unlink(path)
                    count += 1
            except (json.JSONDecodeError, OSError):
                try:
                    os.unlink(path)
                except OSError:
                    pass
    except OSError:
        pass

    if count:
        debug(f"INVALIDATED {count} entries for {resource_type} {resource_id}")
        stats = read_stats()
        stats["invalidations"] = stats.get("invalidations", 0) + count
        write_stats(stats)


def invalidate_all_for_type(resource_type: str) -> None:
    """Broad invalidation: clear all entries when we can't determine the resource id."""
    debug(f"INVALIDATE ALL (mutation on {resource_type})")
    clear_cache()


def clear_cache() -> None:
    """Remove all cached entries."""
    try:
        for name in os.listdir(CACHE_DIR):
            if name.startswith("_"):
                continue
            path = os.path.join(CACHE_DIR, name)
            os.unlink(path)
    except OSError:
        pass


# ─── Command Analysis ────────────────────────────────────────────────────────

def parse_gh_args(args: list[str]) -> dict:
    """Parse gh command arguments to determine cacheability.

    Returns dict with:
        resource_type: "issue", "pr", "api", etc.
        subcommand: "view", "list", "edit", etc.
        resource_id: The issue/PR number if present
        cacheable: Whether this command can be cached
        ttl: TTL to use for this command
    """
    result = {
        "resource_type": None,
        "subcommand": None,
        "resource_id": None,
        "cacheable": False,
        "ttl": DEFAULT_TTL,
    }

    if not args:
        return result

    # First non-flag arg is the resource type
    resource_type = args[0]
    result["resource_type"] = resource_type

    if resource_type in PASSTHROUGH_COMMANDS:
        return result

    # Special case: `gh api` — cache GET requests
    if resource_type == "api":
        result["subcommand"] = "api"
        # Check for method flags that indicate mutation
        is_mutation = False
        for i, arg in enumerate(args):
            if arg in ("-X", "--method") and i + 1 < len(args):
                method = args[i + 1].upper()
                if method != "GET":
                    is_mutation = True
                    break
            if arg == "-f" or arg == "--field":
                # POST with fields
                is_mutation = True
                break
        if not is_mutation:
            result["cacheable"] = True
            result["ttl"] = TTL_BY_COMMAND.get(("api",), DEFAULT_TTL)
        return result

    # Second arg is the subcommand
    if len(args) < 2:
        return result

    subcommand = args[1]
    result["subcommand"] = subcommand

    # Extract resource ID (first non-flag arg after subcommand)
    for arg in args[2:]:
        if not arg.startswith("-"):
            result["resource_id"] = arg
            break

    # Determine TTL
    ttl_key = (resource_type, subcommand)
    result["ttl"] = TTL_BY_COMMAND.get(ttl_key, DEFAULT_TTL)

    # Determine cacheability
    if subcommand in CACHEABLE_SUBCOMMANDS:
        result["cacheable"] = True
    elif subcommand in MUTATION_SUBCOMMANDS:
        result["cacheable"] = False

    return result


# ─── Main ────────────────────────────────────────────────────────────────────

def run_gh(args: list[str]) -> tuple[str, str, int]:
    """Execute the real gh command and return (stdout, stderr, returncode)."""
    proc = subprocess.run(
        [os.environ.get("LOOM_GH_BIN") or "gh"] + args,
        capture_output=True,
        text=True,
    )
    return proc.stdout, proc.stderr, proc.returncode


# ─── ETag/REST cached listing (#5056) ────────────────────────────────────────

def locate_loom_daemon() -> str | None:
    """Resolve a loom-daemon binary for the ETag-cached `forge … list --cached`
    path. Mirrors lib/locate-daemon-bin.sh's common cases: $LOOM_DAEMON_BIN, then
    PATH, then the machine-level install. Returns None when none is executable —
    the caller then degrades to plain `gh` (the "daemon unreachable" fallback)."""
    bin_env = os.environ.get("LOOM_DAEMON_BIN", "")
    if bin_env and os.access(bin_env, os.X_OK):
        return bin_env
    from shutil import which
    found = which("loom-daemon")
    if found:
        return found
    machine = os.path.join(
        os.environ.get("LOOM_DAEMON_BIN_DIR", os.path.expanduser("~/.local/bin")),
        "loom-daemon",
    )
    if os.access(machine, os.X_OK):
        return machine
    return None


def try_etag_list(resource_type: str, args: list[str]) -> tuple[str, int] | None:
    """Route a cacheable `issue list` / `pr list` through loom-daemon's
    disk-persistent ETag REST cache (#5056).

    A validated 304 costs **zero** rate-limit units — and, unlike the local TTL
    cache below, is never stale (a 304 is positive proof nothing changed). So
    this is tried *before* the TTL cache for the label/state listing shape.

    Returns (stdout, 0) when the daemon served the query, or None to fall back
    to the normal path (binary absent → daemon unreachable; or the daemon
    declined an uncacheable shape — e.g. `--search head:…`, a table `list` with
    no `--json`, a PR-only field — with a non-zero exit)."""
    if etag_layer_disabled():
        return None
    daemon = locate_loom_daemon()
    if not daemon:
        return None
    # args == [resource_type, "list", <rest…>]; forward everything after "list".
    rest = args[2:]
    cmd = [daemon, "forge", resource_type, "list", "--cached", *rest]
    try:
        proc = subprocess.run(cmd, capture_output=True, text=True)
    except OSError:
        return None
    if proc.returncode == 0:
        return proc.stdout, 0
    return None


def etag_layer_disabled(*extra_vars: str) -> bool:
    """LOOM_ETAG_LIST_DISABLE=1 turns off the whole ETag layer (list + view);
    extra_vars are per-shape kill switches (e.g. LOOM_ETAG_VIEW_DISABLE)."""
    return any(os.environ.get(v, "") == "1" for v in ("LOOM_ETAG_LIST_DISABLE", *extra_vars))


def try_etag_view(resource_type: str, args: list[str]) -> tuple[str, int] | None:
    """Route `issue view N --json …` / `pr view N --json …` through loom-daemon's
    conditional-request (ETag/304) single-object read (#9254).

    Every read is revalidated with If-None-Match, so — unlike the TTL cache
    below — it is never stale, and an unchanged object costs a free 304. Tried
    before the TTL cache; returns None (fall back) when the binary is absent or
    the daemon declines the shape (exit 3: no --json, an unsupported field such
    as author/mergeable/statusCheckRollup, --comments/--web, a non-numeric
    selector, an issue number that is really a PR, …)."""
    if etag_layer_disabled("LOOM_ETAG_VIEW_DISABLE"):
        return None
    daemon = locate_loom_daemon()
    if not daemon:
        return None
    # args == [resource_type, "view", <rest…>]; forward everything after "view".
    cmd = [daemon, "forge", resource_type, "view", "--cached", *args[2:]]
    try:
        proc = subprocess.run(cmd, capture_output=True, text=True)
    except OSError:
        return None
    if proc.returncode == 0:
        return proc.stdout, 0
    return None


def invalidate_etag_views(resource_id: str | None) -> None:
    """Write-through rule (ADR-0021 amendment, #9254): after a wrapped mutation,
    drop the daemon's ETag view entries for that number so the next read is
    unconditional (guards against a replica-lag 304 right after our own write).
    A non-numeric/absent id drops every view entry. Best-effort."""
    if etag_layer_disabled("LOOM_ETAG_VIEW_DISABLE"):
        return
    daemon = locate_loom_daemon()
    if not daemon:
        return
    number = None
    if resource_id:
        tail = resource_id.rstrip("/").rsplit("/", 1)[-1].lstrip("#")
        if tail.isdigit():
            number = tail
    cmd = [daemon, "forge", "issue", "view", "--cached", "--invalidate"]
    if number:
        cmd.append(number)
    try:
        subprocess.run(cmd, capture_output=True, text=True)
    except OSError:
        pass


def main() -> int:
    args = sys.argv[1:]
    debug(f"REPO_ID={REPO_ID!r} CACHE_DIR={CACHE_DIR}")

    # Handle meta-commands
    if "--clear-cache" in args:
        clear_cache()
        print("Cache cleared.", file=sys.stderr)
        return 0

    if "--cache-stats" in args:
        stats = read_stats()
        total = stats["hits"] + stats["misses"]
        rate = (stats["hits"] / total * 100) if total > 0 else 0
        print(f"Hits: {stats['hits']}", file=sys.stderr)
        print(f"Misses: {stats['misses']}", file=sys.stderr)
        print(f"Bypasses: {stats['bypasses']}", file=sys.stderr)
        print(f"Invalidations: {stats['invalidations']}", file=sys.stderr)
        print(f"ETag (REST conditional list/view reads): {stats.get('etag', 0)}", file=sys.stderr)
        print(f"Hit rate: {rate:.1f}%", file=sys.stderr)
        return 0

    # Handle --no-cache flag
    no_cache = False
    if "--no-cache" in args:
        args.remove("--no-cache")
        no_cache = True

    # Disabled or passthrough
    if CACHE_DISABLED or no_cache or not args:
        stdout, stderr, rc = run_gh(args)
        sys.stdout.write(stdout)
        sys.stderr.write(stderr)
        if no_cache:
            increment_stat("bypasses")
        record_outcome("bypass", args)
        return rc

    parsed = parse_gh_args(args)
    debug(f"CMD: gh {' '.join(args)}")
    debug(f"PARSED: cacheable={parsed['cacheable']} type={parsed['resource_type']} "
          f"sub={parsed['subcommand']} id={parsed['resource_id']} ttl={parsed['ttl']}")

    # Mutation: bypass cache, invalidate related entries
    if not parsed["cacheable"]:
        stdout, stderr, rc = run_gh(args)
        sys.stdout.write(stdout)
        sys.stderr.write(stderr)

        # Invalidate on successful mutations
        if rc == 0 and parsed["subcommand"] in MUTATION_SUBCOMMANDS:
            if parsed["resource_id"]:
                invalidate_for_resource(parsed["resource_type"], parsed["resource_id"])
            else:
                invalidate_all_for_type(parsed["resource_type"] or "unknown")
            if parsed["resource_type"] in ("issue", "pr"):
                invalidate_etag_views(parsed["resource_id"])

        increment_stat("bypasses")
        record_outcome("bypass", args)
        return rc

    # ETag/REST cached listing (#5056): `issue list` / `pr list` can be served
    # for free on a validated 304 through loom-daemon's disk-persistent cache —
    # a separate rate-limit pool (REST, not GraphQL) AND zero cost when nothing
    # changed. Tried before the local TTL cache because a 304 is both free and
    # never stale. On decline / absent binary it falls through to the TTL path.
    if parsed["subcommand"] == "list" and parsed["resource_type"] in ("issue", "pr"):
        etag_result = try_etag_list(parsed["resource_type"], args)
        if etag_result is not None:
            etag_stdout, etag_rc = etag_result
            sys.stdout.write(etag_stdout)
            increment_stat("etag")
            record_outcome("revalidated", args)
            return etag_rc

    # Conditional single-object read (#9254): `issue view` / `pr view` are
    # revalidated with If-None-Match (free 304, never stale) before the TTL
    # cache; a decline falls through to the TTL path below unchanged.
    if parsed["subcommand"] == "view" and parsed["resource_type"] in ("issue", "pr"):
        etag_result = try_etag_view(parsed["resource_type"], args)
        if etag_result is not None:
            etag_stdout, etag_rc = etag_result
            sys.stdout.write(etag_stdout)
            increment_stat("etag")
            record_outcome("revalidated", args)
            return etag_rc

    # Read-only: check cache
    key = cache_key(args)
    cached_stdout, cached_rc = cache_get(key)

    if cached_stdout is not None:
        debug(f"HIT key={key}")
        sys.stdout.write(cached_stdout)
        increment_stat("hits")
        record_outcome("hit", args)
        return cached_rc

    # Cache miss: run real command
    debug(f"MISS key={key}")
    stdout, stderr, rc = run_gh(args)
    sys.stdout.write(stdout)
    sys.stderr.write(stderr)

    # Only cache successful responses
    if rc == 0:
        cache_put(key, stdout, rc, parsed["ttl"], args=args)

    increment_stat("misses")
    record_outcome("miss", args)
    return rc


if __name__ == "__main__":
    sys.exit(main())
