#!/usr/bin/env python3
from __future__ import annotations

import argparse
import fcntl
import json
import os
import subprocess
import sys
import time
import uuid
from contextlib import contextmanager
from datetime import datetime, timezone
from pathlib import Path
from typing import Any

EXIT_OK = 0
EXIT_ERROR = 1
EXIT_BUSY = 2

DEFAULT_TIMEOUT = 600
DEFAULT_TTL = 900
DEFAULT_POLL = 5
DEFAULT_PREFER = "shutdown"


def pool_home() -> Path:
    return Path(os.environ.get("AGENT_SIM_POOL_HOME", Path.home() / ".agent-sim-pool"))


def resolve_holder_pid(explicit: int | None) -> int:
    if explicit is not None and explicit > 0:
        return explicit
    override = os.environ.get("SIM_POOL_LEASE_PID", "").strip()
    if override:
        return int(override)
    return os.getppid()


def now_iso() -> str:
    return datetime.now(timezone.utc).replace(microsecond=0).isoformat()


def parse_iso(value: str) -> datetime:
    return datetime.fromisoformat(value.replace("Z", "+00:00"))


def ensure_layout(home: Path) -> None:
    home.mkdir(parents=True, exist_ok=True)
    (home / "leases").mkdir(parents=True, exist_ok=True)
    lock_path = home / "pool.lock"
    if not lock_path.exists():
        lock_path.touch()


def config_path(home: Path) -> Path:
    return home / "config.json"


def default_config() -> dict[str, Any]:
    return {
        "devices": [],
        "defaults": {
            "timeout_seconds": DEFAULT_TIMEOUT,
            "ttl_seconds": DEFAULT_TTL,
            "poll_seconds": DEFAULT_POLL,
            "prefer": DEFAULT_PREFER,
        },
    }


def load_config(home: Path) -> dict[str, Any]:
    path = config_path(home)
    if not path.is_file():
        return default_config()
    data = json.loads(path.read_text(encoding="utf-8"))
    defaults = default_config()["defaults"]
    defaults.update(data.get("defaults", {}))
    data["defaults"] = defaults
    if "devices" not in data:
        data["devices"] = []
    return data


def save_config(home: Path, config: dict[str, Any]) -> None:
    ensure_layout(home)
    config_path(home).write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8")


def lease_path(home: Path, udid: str) -> Path:
    safe = udid.replace("/", "_")
    return home / "leases" / f"{safe}.json"


@contextmanager
def pool_lock(home: Path):
    ensure_layout(home)
    lock_file = open(home / "pool.lock", "a+")
    try:
        fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX)
        yield
    finally:
        fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN)
        lock_file.close()


def pid_alive(pid: int) -> bool:
    if pid <= 0:
        return False
    try:
        os.kill(pid, 0)
    except OSError:
        return False
    return True


def lease_expired(lease: dict[str, Any], now: datetime | None = None) -> bool:
    now = now or datetime.now(timezone.utc)
    expires_at = lease.get("expires_at")
    if not expires_at:
        return True
    return parse_iso(expires_at) <= now


def lease_corrupt(lease: dict[str, Any], ttl_seconds: int, now: datetime | None = None) -> bool:
    now = now or datetime.now(timezone.utc)
    expires_at = lease.get("expires_at")
    acquired_at = lease.get("acquired_at")
    if not expires_at or not acquired_at:
        return True
    try:
        exp = parse_iso(expires_at)
        acq = parse_iso(acquired_at)
    except ValueError:
        return True
    if exp < acq:
        return True
    if (exp - now).total_seconds() > ttl_seconds * 2:
        return True
    return False


def read_lease(path: Path) -> dict[str, Any] | None:
    if not path.is_file():
        return None
    try:
        return json.loads(path.read_text(encoding="utf-8"))
    except (json.JSONDecodeError, OSError):
        return None


def write_lease(path: Path, lease: dict[str, Any]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    tmp = path.with_suffix(".json.tmp")
    tmp.write_text(json.dumps(lease, indent=2) + "\n", encoding="utf-8")
    tmp.replace(path)


def remove_lease(path: Path) -> bool:
    if path.is_file():
        path.unlink()
        return True
    return False


def simctl_devices() -> dict[str, dict[str, Any]]:
    try:
        proc = subprocess.run(
            ["xcrun", "simctl", "list", "devices", "available", "-j"],
            check=True,
            capture_output=True,
            text=True,
        )
    except (subprocess.CalledProcessError, FileNotFoundError):
        return {}
    data = json.loads(proc.stdout)
    out: dict[str, dict[str, Any]] = {}
    for runtime, devices in data.get("devices", {}).items():
        if "iOS" not in runtime:
            continue
        for device in devices:
            if not device.get("isAvailable", True):
                continue
            name = device.get("name", "")
            udid = device.get("udid", "")
            if not udid or "iPhone" not in name:
                continue
            state = device.get("state", "Shutdown")
            out[udid] = {"name": name, "state": state, "runtime": runtime}
    return out


def discover_iphone_udids(limit: int = 10) -> list[str]:
    devices = simctl_devices()
    preferred = (
        "iPhone 16 Pro",
        "iPhone 16",
        "iPhone 15 Pro",
        "iPhone 15",
        "iPhone SE (3rd generation)",
    )

    def sort_key(item: tuple[str, dict[str, Any]]) -> tuple[int, str]:
        _udid, info = item
        name = info.get("name", "")
        for idx, pref in enumerate(preferred):
            if name == pref:
                return (idx, name)
        return (len(preferred), name)

    ranked = sorted(devices.items(), key=sort_key)
    return [udid for udid, _ in ranked[:limit]]


def init_config(home: Path, force: bool = False) -> dict[str, Any]:
    ensure_layout(home)
    path = config_path(home)
    if path.is_file() and not force:
        return load_config(home)
    config = default_config()
    config["devices"] = discover_iphone_udids()
    save_config(home, config)
    return config


def active_leases(home: Path, config: dict[str, Any]) -> dict[str, dict[str, Any]]:
    ttl = int(config["defaults"].get("ttl_seconds", DEFAULT_TTL))
    leases: dict[str, dict[str, Any]] = {}
    leases_dir = home / "leases"
    if not leases_dir.is_dir():
        return leases
    for path in leases_dir.glob("*.json"):
        lease = read_lease(path)
        if not lease:
            continue
        udid = lease.get("udid", "")
        if not udid:
            continue
        leases[udid] = lease
    return leases


def should_drop_lease(lease: dict[str, Any], config: dict[str, Any], now: datetime | None = None) -> tuple[bool, str]:
    now = now or datetime.now(timezone.utc)
    ttl = int(config["defaults"].get("ttl_seconds", DEFAULT_TTL))
    if lease_corrupt(lease, ttl, now):
        return True, "corrupt"
    pid = int(lease.get("pid", 0))
    if not pid_alive(pid):
        return True, "dead_pid"
    if lease_expired(lease, now):
        return True, "expired"
    return False, ""


def gc(home: Path, config: dict[str, Any]) -> list[str]:
    dropped: list[str] = []
    leases_dir = home / "leases"
    if not leases_dir.is_dir():
        return dropped
    for path in sorted(leases_dir.glob("*.json")):
        lease = read_lease(path)
        if not lease:
            path.unlink(missing_ok=True)
            dropped.append(path.stem)
            continue
        drop, reason = should_drop_lease(lease, config)
        if drop:
            udid = lease.get("udid", path.stem)
            remove_lease(path)
            dropped.append(f"{udid}:{reason}")
    return dropped


def leased_udids(home: Path, config: dict[str, Any]) -> set[str]:
    gc(home, config)
    out: set[str] = set()
    for path in (home / "leases").glob("*.json"):
        lease = read_lease(path)
        if not lease:
            continue
        drop, _ = should_drop_lease(lease, config)
        if drop:
            continue
        udid = lease.get("udid")
        if udid:
            out.add(udid)
    return out


def device_sort_key(
    udid: str,
    sim_info: dict[str, dict[str, Any]],
    prefer: str,
) -> tuple[int, int, str]:
    info = sim_info.get(udid, {})
    name = info.get("name", udid)
    state = info.get("state", "Shutdown")
    shutdown_rank = 0 if state == "Shutdown" else 1
    if prefer == "booted":
        shutdown_rank = 0 if state != "Shutdown" else 1
    return (shutdown_rank, 0, name)


def pick_free_udid(
    home: Path,
    config: dict[str, Any],
    prefer_udid: str | None,
    prefer: str,
) -> str | None:
    devices = config.get("devices", [])
    if not devices:
        return None
    leased = leased_udids(home, config)
    sim_info = simctl_devices()
    free = [udid for udid in devices if udid not in leased]
    if not free:
        return None
    if prefer_udid and prefer_udid in free:
        return prefer_udid
    free.sort(key=lambda u: device_sort_key(u, sim_info, prefer))
    return free[0]


def find_lease_by_id(home: Path, lease_id: str) -> tuple[Path | None, dict[str, Any] | None]:
    for path in (home / "leases").glob("*.json"):
        lease = read_lease(path)
        if lease and lease.get("lease_id") == lease_id:
            return path, lease
    return None, None


def cmd_status(home: Path) -> int:
    config = load_config(home)
    gc(home, config)
    devices = config.get("devices", [])
    sim_info = simctl_devices()
    now = datetime.now(timezone.utc)
    print(f"pool_home={home}")
    print(f"devices_whitelisted={len(devices)}")
    for udid in devices:
        path = lease_path(home, udid)
        lease = read_lease(path)
        name = sim_info.get(udid, {}).get("name", "unknown")
        state = sim_info.get(udid, {}).get("state", "unknown")
        if not lease:
            print(f"FREE\t{udid}\t{name}\t{state}")
            continue
        drop, reason = should_drop_lease(lease, config, now)
        if drop:
            print(f"STALE\t{udid}\t{name}\t{state}\treason={reason}")
            continue
        owner = lease.get("owner", "")
        project = lease.get("project", "")
        pid = lease.get("pid", "")
        expires = lease.get("expires_at", "")
        lease_id = lease.get("lease_id", "")
        print(
            f"LEASED\t{udid}\t{name}\t{state}\t"
            f"lease_id={lease_id}\towner={owner}\tproject={project}\tpid={pid}\texpires_at={expires}"
        )
    return EXIT_OK


def cmd_acquire(args: argparse.Namespace, home: Path) -> int:
    config = load_config(home)
    if not config.get("devices"):
        init_config(home)
        config = load_config(home)
    defaults = config["defaults"]
    timeout = args.timeout if args.timeout is not None else int(defaults.get("timeout_seconds", DEFAULT_TIMEOUT))
    ttl = args.ttl if args.ttl is not None else int(defaults.get("ttl_seconds", DEFAULT_TTL))
    poll = int(defaults.get("poll_seconds", DEFAULT_POLL))
    prefer = defaults.get("prefer", DEFAULT_PREFER)
    deadline = time.monotonic() + timeout
    owner = args.owner or f"agent-{os.getpid()}"
    while True:
        with pool_lock(home):
            gc(home, config)
            udid = pick_free_udid(home, config, args.prefer_udid, prefer)
            if udid:
                lease_id = str(uuid.uuid4())
                acquired_at = now_iso()
                expires_at = (
                    datetime.now(timezone.utc).timestamp() + ttl
                )
                expires_iso = datetime.fromtimestamp(expires_at, timezone.utc).replace(microsecond=0).isoformat()
                sim_info = simctl_devices()
                name = sim_info.get(udid, {}).get("name", "unknown")
                lease = {
                    "lease_id": lease_id,
                    "udid": udid,
                    "owner": owner,
                    "pid": resolve_holder_pid(args.holder_pid),
                    "project": args.project or "",
                    "worktree": args.worktree or "",
                    "session": args.session or "",
                    "purpose": args.purpose or "qa",
                    "acquired_at": acquired_at,
                    "expires_at": expires_iso,
                }
                write_lease(lease_path(home, udid), lease)
                print(f"LEASE_ID={lease_id}")
                print(f"UDID={udid}")
                print(f"SIMULATOR_NAME={name}")
                print(f"EXPIRES_AT={expires_iso}")
                return EXIT_OK
        if time.monotonic() >= deadline:
            print("SIM_POOL_BUSY", file=sys.stderr)
            return EXIT_BUSY
        time.sleep(poll)


def cmd_renew(args: argparse.Namespace, home: Path) -> int:
    config = load_config(home)
    ttl = args.ttl if args.ttl is not None else int(config["defaults"].get("ttl_seconds", DEFAULT_TTL))
    with pool_lock(home):
        path, lease = find_lease_by_id(home, args.lease)
        if not lease or not path:
            print("LEASE_GONE", file=sys.stderr)
            return EXIT_ERROR
        drop, reason = should_drop_lease(lease, config)
        if drop:
            remove_lease(path)
            print(f"LEASE_GONE reason={reason}", file=sys.stderr)
            return EXIT_ERROR
        expires_iso = datetime.fromtimestamp(
            datetime.now(timezone.utc).timestamp() + ttl,
            timezone.utc,
        ).replace(microsecond=0).isoformat()
        lease["expires_at"] = expires_iso
        write_lease(path, lease)
        print(f"LEASE_ID={lease['lease_id']}")
        print(f"EXPIRES_AT={expires_iso}")
        return EXIT_OK


def cmd_release(args: argparse.Namespace, home: Path) -> int:
    with pool_lock(home):
        if args.lease:
            path, lease = find_lease_by_id(home, args.lease)
            if path and lease:
                remove_lease(path)
        elif args.udid:
            remove_lease(lease_path(home, args.udid))
        else:
            print("release requires --lease or --udid", file=sys.stderr)
            return EXIT_ERROR
    return EXIT_OK


def cmd_force_release(args: argparse.Namespace, home: Path) -> int:
    with pool_lock(home):
        path = lease_path(home, args.udid)
        if remove_lease(path):
            reason = args.reason or "force-release"
            print(f"FORCE_RELEASED udid={args.udid} reason={reason}")
        else:
            print(f"NO_LEASE udid={args.udid}")
    return EXIT_OK


def cmd_register(args: argparse.Namespace, home: Path) -> int:
    ensure_layout(home)
    config = load_config(home)
    devices = list(config.get("devices", []))
    if args.udid not in devices:
        devices.append(args.udid)
    config["devices"] = devices
    save_config(home, config)
    print(f"REGISTERED {args.udid}")
    return EXIT_OK


def cmd_unregister(args: argparse.Namespace, home: Path) -> int:
    config = load_config(home)
    devices = [d for d in config.get("devices", []) if d != args.udid]
    config["devices"] = devices
    save_config(home, config)
    remove_lease(lease_path(home, args.udid))
    print(f"UNREGISTERED {args.udid}")
    return EXIT_OK


def agent_device_claim(udid: str) -> str | None:
    claims_dir = Path.home() / ".agent-device" / "device-claims"
    if not claims_dir.is_dir():
        return None
    for path in claims_dir.glob("*.json"):
        try:
            data = json.loads(path.read_text(encoding="utf-8"))
        except (json.JSONDecodeError, OSError):
            continue
        device = data.get("device", {})
        if device.get("id") == udid:
            return data.get("session", path.stem)
    return None


def cmd_doctor(home: Path) -> int:
    config = load_config(home)
    if not config.get("devices"):
        print("config.devices is empty — run: sim-pool init")
    sim_info = simctl_devices()
    print("=== sim-pool doctor ===")
    print(f"pool_home: {home}")
    print(f"whitelisted: {len(config.get('devices', []))}")
    for udid in config.get("devices", []):
        known = udid in sim_info
        name = sim_info.get(udid, {}).get("name", "MISSING")
        lease = read_lease(lease_path(home, udid))
        claim = agent_device_claim(udid)
        status = "ok" if known else "missing_from_simctl"
        print(f"  {udid}  {name}  [{status}]")
        if lease and not should_drop_lease(lease, config)[0]:
            print(f"    lease: {lease.get('lease_id')} owner={lease.get('owner')} expires={lease.get('expires_at')}")
        if claim and not lease:
            print(f"    warn: agent-device claim session={claim} but no sim-pool lease")
    print("")
    print("Agent recipe:")
    print("  sim-pool acquire --holder-pid $$ --owner <unique> --project <repo> --worktree \"$PWD\" --session <name>")
    print("  export DEVICE_ID=<UDID from output>")
    print("  agent-device open ... --udid $DEVICE_ID --session <name>")
    print("  sim-pool renew --lease <id>   # every ~5 min on long QA")
    print("  agent-device close --session <name>")
    print("  sim-pool release --lease <id>")
    print("")
    print("Orphan recovery: TTL + dead-pid GC; release is polite not required.")
    print("On SIM_POOL_BUSY (exit 2): report inconclusive; do not create simulators.")
    return EXIT_OK


def cmd_init(home: Path, force: bool) -> int:
    config = init_config(home, force=force)
    print(f"INIT pool_home={home} devices={len(config.get('devices', []))}")
    for udid in config.get("devices", []):
        print(f"  {udid}")
    return EXIT_OK


def cmd_gc(home: Path) -> int:
    config = load_config(home)
    with pool_lock(home):
        dropped = gc(home, config)
    for item in dropped:
        print(f"DROPPED {item}")
    if not dropped:
        print("GC nothing to drop")
    return EXIT_OK


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(prog="sim-pool")
    sub = parser.add_subparsers(dest="command", required=True)

    sub.add_parser("status", help="Show pool and lease state")
    sub.add_parser("gc", help="Drop dead-pid and expired leases")
    sub.add_parser("doctor", help="Validate pool and print agent recipe")

    p_init = sub.add_parser("init", help="Discover iPhone sims into whitelist (no create)")
    p_init.add_argument("--force", action="store_true")

    p_acquire = sub.add_parser("acquire", help="Block until a simulator lease is granted")
    p_acquire.add_argument("--owner", default="")
    p_acquire.add_argument("--project", default="")
    p_acquire.add_argument("--worktree", default="")
    p_acquire.add_argument("--purpose", default="qa")
    p_acquire.add_argument("--session", default="")
    p_acquire.add_argument("--prefer-udid", default="")
    p_acquire.add_argument("--timeout", type=int, default=None)
    p_acquire.add_argument("--ttl", type=int, default=None)
    p_acquire.add_argument("--holder-pid", type=int, default=None)

    p_renew = sub.add_parser("renew", help="Extend lease TTL")
    p_renew.add_argument("--lease", required=True)
    p_renew.add_argument("--ttl", type=int, default=None)

    p_release = sub.add_parser("release", help="Drop a lease")
    p_release.add_argument("--lease", default="")
    p_release.add_argument("--udid", default="")

    p_force = sub.add_parser("force-release", help="Drop lease without owner check (user-approved)")
    p_force.add_argument("--udid", required=True)
    p_force.add_argument("--reason", default="")

    p_reg = sub.add_parser("register", help="Add UDID to whitelist")
    p_reg.add_argument("--udid", required=True)

    p_unreg = sub.add_parser("unregister", help="Remove UDID from whitelist")
    p_unreg.add_argument("--udid", required=True)

    return parser


def main() -> int:
    parser = build_parser()
    args = parser.parse_args()
    home = pool_home()

    if args.command == "status":
        return cmd_status(home)
    if args.command == "gc":
        return cmd_gc(home)
    if args.command == "doctor":
        return cmd_doctor(home)
    if args.command == "init":
        return cmd_init(home, force=args.force)
    if args.command == "acquire":
        return cmd_acquire(args, home)
    if args.command == "renew":
        return cmd_renew(args, home)
    if args.command == "release":
        return cmd_release(args, home)
    if args.command == "force-release":
        return cmd_force_release(args, home)
    if args.command == "register":
        return cmd_register(args, home)
    if args.command == "unregister":
        return cmd_unregister(args, home)
    parser.print_help()
    return EXIT_ERROR


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