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

import argparse
import json
import shlex
import subprocess
import sys
import time
from pathlib import Path
from typing import Any

SCRIPT_DIR = Path(__file__).resolve().parent
REGISTRY_PATH = SCRIPT_DIR / "services.registry.json"
GLOBAL_CONF_PATH = SCRIPT_DIR / "ios-mac-services.conf"
PROJECT_CONF_DIR = SCRIPT_DIR / "services.d"
GOVERNANCECTL_PATH = SCRIPT_DIR / "governancectl.py"
GENERATED_HEADER = "# GENERATED BY public/dexrelay-runtime/servicectl.py. DO NOT EDIT DIRECTLY."


def run_shell(command: str, *, timeout: int = 30) -> dict[str, Any]:
    completed = subprocess.run(
        ["/bin/zsh", "-lc", command],
        capture_output=True,
        text=True,
        timeout=timeout,
    )
    return {
        "exitCode": completed.returncode,
        "stdout": completed.stdout,
        "stderr": completed.stderr,
        "command": command,
    }


def print_json(payload: dict[str, Any]) -> None:
    print(json.dumps(payload, indent=2, sort_keys=False))


def slugify(value: str) -> str:
    out: list[str] = []
    prev_dash = False
    for ch in value.lower():
        if ch.isalnum():
            out.append(ch)
            prev_dash = False
        elif not prev_dash:
            out.append("-")
            prev_dash = True
    return "".join(out).strip("-") or "project"


def load_registry(path: Path = REGISTRY_PATH) -> dict[str, Any]:
    if not path.exists():
        return {
            "version": 1,
            "portPolicy": {"min": 8000, "max": 8999, "reserved": [4500, 4610, 4615, 4616, 4620]},
            "services": [],
        }
    data = json.loads(path.read_text(encoding="utf-8"))
    if not isinstance(data, dict):
        raise ValueError("registry root must be an object")
    data.setdefault("version", 1)
    data.setdefault("portPolicy", {})
    data.setdefault("services", [])
    if not isinstance(data["services"], list):
        raise ValueError("registry services must be a list")
    return data


def services(data: dict[str, Any]) -> list[dict[str, Any]]:
    return [item for item in data.get("services", []) if isinstance(item, dict)]


def validate_registry(data: dict[str, Any]) -> tuple[list[str], list[str]]:
    errors: list[str] = []
    warnings: list[str] = []
    seen_ids: set[str] = set()
    claimed_ports: dict[int, str] = {}

    for idx, svc in enumerate(services(data)):
        service_id = str(svc.get("id", "")).strip()
        name = str(svc.get("name", "")).strip()
        start_command = str(svc.get("startCommand", "")).strip()
        project_path = str(svc.get("projectPath", "")).strip()

        if not service_id:
            errors.append(f"service[{idx}] missing id")
            continue
        if service_id in seen_ids:
            errors.append(f"duplicate service id: {service_id}")
        seen_ids.add(service_id)

        if not name:
            errors.append(f"service '{service_id}' missing name")
        if not start_command:
            errors.append(f"service '{service_id}' missing startCommand")
        if not project_path:
            warnings.append(f"service '{service_id}' has empty projectPath; treating as global")

        raw_ports = svc.get("ports", [])
        ports = raw_ports if isinstance(raw_ports, list) else []
        for port in ports:
            if not isinstance(port, int):
                errors.append(f"service '{service_id}' has non-integer port: {port}")
                continue
            if port < 1 or port > 65535:
                errors.append(f"service '{service_id}' has invalid port {port}")
                continue
            other = claimed_ports.get(port)
            if other and other != service_id:
                errors.append(f"port {port} is claimed by both '{other}' and '{service_id}'")
            claimed_ports[port] = service_id

    policy = data.get("portPolicy", {})
    if isinstance(policy, dict):
        for port in policy.get("reserved", []) or []:
            if isinstance(port, int) and port in claimed_ports:
                errors.append(f"reserved port {port} is also claimed by service '{claimed_ports[port]}'")

    return errors, warnings


def port_owner(port: int) -> dict[str, Any] | None:
    result = run_shell(f"lsof -nP -iTCP:{port} -sTCP:LISTEN -Fpcn", timeout=10)
    if result["exitCode"] != 0 or not result["stdout"].strip():
        return None
    pid = None
    command = ""
    for line in result["stdout"].splitlines():
        if line.startswith("p") and line[1:].isdigit():
            pid = int(line[1:])
        elif line.startswith("c"):
            command = line[1:]
    if pid is None:
        return None
    return {"pid": pid, "command": command}


def write_conf(path: Path, lines: list[str]) -> None:
    content = [GENERATED_HEADER, ""]
    content.extend(lines)
    content.append("")
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text("\n".join(content), encoding="utf-8")


def csv_ports(ports: list[int]) -> str:
    return ",".join(str(port) for port in ports)


def sync_conf(data: dict[str, Any]) -> list[Path]:
    grouped: dict[str, list[dict[str, Any]]] = {}
    for svc in services(data):
        project_path = str(svc.get("projectPath", "")).strip() or "*"
        grouped.setdefault(project_path, []).append(svc)

    written: list[Path] = []

    global_lines: list[str] = []
    for svc in grouped.get("*", []):
        ports = [port for port in svc.get("ports", []) if isinstance(port, int)]
        health = str(svc.get("healthCheck", "")).strip()
        start = str(svc.get("startCommand", "")).strip()
        name = str(svc.get("name", "")).strip()
        if not name or not start:
            continue
        line = f"{name} | {health} | {start}"
        if ports:
            line += f" | {csv_ports(ports)}"
        global_lines.append(line)

    write_conf(GLOBAL_CONF_PATH, global_lines)
    written.append(GLOBAL_CONF_PATH)

    expected: set[Path] = set()
    for project_path, items in grouped.items():
        if project_path == "*":
            continue
        conf_path = PROJECT_CONF_DIR / f"{slugify(Path(project_path).name)}.conf"
        expected.add(conf_path)
        lines = [f"# Active when current project path basename slug is: {slugify(Path(project_path).name)}"]
        for svc in items:
            ports = [port for port in svc.get("ports", []) if isinstance(port, int)]
            health = str(svc.get("healthCheck", "")).strip()
            start = str(svc.get("startCommand", "")).strip()
            name = str(svc.get("name", "")).strip()
            if not name or not start:
                continue
            line = f"{name} | {health} | {start}"
            if ports:
                line += f" | {csv_ports(ports)}"
            lines.append(line)
        write_conf(conf_path, lines)
        written.append(conf_path)

    PROJECT_CONF_DIR.mkdir(parents=True, exist_ok=True)
    for candidate in PROJECT_CONF_DIR.glob("*.conf"):
        if candidate in expected:
            continue
        try:
            content = candidate.read_text(encoding="utf-8")
        except OSError:
            continue
        if GENERATED_HEADER in content:
            candidate.unlink()

    return written


def check_health(command: str) -> bool | None:
    trimmed = command.strip()
    if not trimmed:
        return None
    result = run_shell(trimmed, timeout=15)
    return result["exitCode"] == 0


def service_status(svc: dict[str, Any]) -> dict[str, Any]:
    ports = [port for port in svc.get("ports", []) if isinstance(port, int)]
    health = str(svc.get("healthCheck", "")).strip()
    raw_exposure = svc.get("exposure", {})
    exposure = raw_exposure if isinstance(raw_exposure, dict) else {}
    owners = [owner for port in ports if (owner := port_owner(port))]
    open_path = str(svc.get("openPath", "")).strip()
    return {
        "id": str(svc.get("id", "")).strip(),
        "name": str(svc.get("name", "")).strip(),
        "projectPath": str(svc.get("projectPath", "")).strip(),
        "ports": ports,
        "tags": [str(tag) for tag in svc.get("tags", []) if str(tag).strip()],
        "healthCheck": health,
        "healthOk": check_health(health),
        "running": bool(ports) and len(owners) == len(ports) if ports else False,
        "portOwners": owners,
        "startCommand": str(svc.get("startCommand", "")).strip(),
        "stopCommand": str(svc.get("stopCommand", "")).strip(),
        "restartCommand": str(svc.get("restartCommand", "")).strip(),
        "openPath": open_path,
        "exposure": {
            "mode": str(exposure.get("mode", "")).strip(),
            "path": str(exposure.get("path", "")).strip(),
            "target": str(exposure.get("target", "")).strip(),
        },
    }


def wait_for_ports(ports: list[int], *, should_listen: bool, attempts: int = 40) -> bool:
    if not ports:
        return True
    for _ in range(attempts):
        owners = [port_owner(port) for port in ports]
        all_match = all(owner is not None for owner in owners) if should_listen else all(owner is None for owner in owners)
        if all_match:
            return True
        time.sleep(0.5)
    return False


def find_service(service_id: str, data: dict[str, Any]) -> dict[str, Any]:
    for svc in services(data):
        if str(svc.get("id", "")).strip() == service_id:
            return svc
    raise KeyError(service_id)


def service_action(service_id: str, action: str, data: dict[str, Any]) -> dict[str, Any]:
    svc = find_service(service_id, data)
    command_key = {"start": "startCommand", "stop": "stopCommand", "restart": "restartCommand"}[action]
    command = str(svc.get(command_key, "")).strip()
    if not command and action == "restart":
        stop_command = str(svc.get("stopCommand", "")).strip()
        start_command = str(svc.get("startCommand", "")).strip()
        if stop_command and start_command:
            command = f"{stop_command} && {start_command}"
    if not command:
        return {"ok": False, "error": f"service '{service_id}' has no {command_key}"}

    result = run_shell(command, timeout=120)
    ports = [port for port in svc.get("ports", []) if isinstance(port, int)]
    if action == "start":
        wait_for_ports(ports, should_listen=True)
    elif action == "stop":
        wait_for_ports(ports, should_listen=False)
    elif action == "restart":
        wait_for_ports(ports, should_listen=True)

    payload = {
        "ok": result["exitCode"] == 0,
        "action": action,
        "service": service_status(svc),
        "stdout": result["stdout"],
        "stderr": result["stderr"],
        "exitCode": result["exitCode"],
    }

    if GOVENANCE_AVAILABLE:
        run_shell(f"python3 {shlex.quote(str(GOVERNANCECTL_PATH))} summary --json", timeout=60)

    return payload


GOVENANCE_AVAILABLE = GOVERNANCECTL_PATH.exists()


def cmd_list(args: argparse.Namespace) -> int:
    data = load_registry()
    payload = {
        "ok": True,
        "registryPath": str(REGISTRY_PATH),
        "services": [service_status(svc) for svc in services(data)],
    }
    print_json(payload) if args.json else print(f"services={len(payload['services'])}")
    return 0


def cmd_validate(args: argparse.Namespace) -> int:
    data = load_registry()
    errors, warnings = validate_registry(data)
    payload = {"ok": not errors, "errors": errors, "warnings": warnings, "registryPath": str(REGISTRY_PATH)}
    print_json(payload) if args.json else print("ok" if payload["ok"] else "invalid")
    return 0 if not errors else 1


def cmd_sync_conf(args: argparse.Namespace) -> int:
    data = load_registry()
    written = [str(path) for path in sync_conf(data)]
    payload = {"ok": True, "written": written}
    print_json(payload) if args.json else print("\n".join(written))
    return 0


def cmd_start_stop_restart(args: argparse.Namespace) -> int:
    data = load_registry()
    payload = service_action(args.service_id, args.command_name, data)
    print_json(payload) if args.json else print(json.dumps(payload, indent=2))
    return 0 if payload.get("ok") else 1


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description="Manage DexRelay runtime services")
    sub = parser.add_subparsers(dest="command_name", required=True)

    list_parser = sub.add_parser("list")
    list_parser.add_argument("--json", action="store_true")
    list_parser.set_defaults(func=cmd_list)

    validate_parser = sub.add_parser("validate")
    validate_parser.add_argument("--json", action="store_true")
    validate_parser.set_defaults(func=cmd_validate)

    sync_parser = sub.add_parser("sync-conf")
    sync_parser.add_argument("--json", action="store_true")
    sync_parser.set_defaults(func=cmd_sync_conf)

    for name in ("start", "stop", "restart"):
        action_parser = sub.add_parser(name)
        action_parser.add_argument("service_id")
        action_parser.add_argument("--json", action="store_true")
        action_parser.set_defaults(func=cmd_start_stop_restart, command_name=name)

    return parser


def main(argv: list[str] | None = None) -> int:
    parser = build_parser()
    args = parser.parse_args(argv)
    return args.func(args)


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