#!/usr/bin/env python3
"""Usage-aware gaud role preflight.

Reads gaud JSONL config, optionally reads Back2Vibing-style usage snapshots, and
prints a role/agent readiness summary before gaud launches tmux panes.
"""

from __future__ import annotations

import argparse
import json
import os
import sys
import time
from pathlib import Path
from typing import Any

KNOWN_PROVIDERS = {
    "claude": "Claude Code",
    "codex": "Codex",
    "gemini": "Gemini CLI",
    "opencode": "OpenCode",
    "cursor": "Cursor",
    "antigravity": "Antigravity",
    "windsurf": "Windsurf",
}

PROVIDER_ALIASES = {
    "claude-code": "claude",
    "claude_code": "claude",
    "claude": "claude",
    "codex": "codex",
    "openai-codex": "codex",
    "openai_codex": "codex",
    "gemini": "gemini",
    "gemini-cli": "gemini",
    "gemini_cli": "gemini",
    "opencode": "opencode",
    "open-code": "opencode",
    "open_code": "opencode",
    "cursor": "cursor",
    "antigravity": "antigravity",
    "windsurf": "windsurf",
}

STATUS_RANK = {
    "ready": 0,
    "unknown": 1,
    "rate-limited": 2,
    "quota-blocked": 3,
    "auth-blocked": 4,
    "unavailable": 5,
}
MAX_USAGE_CACHE_AGE_MS = 30 * 60 * 1000


def load_last_jsonl(path: Path) -> dict[str, Any]:
    if not path.exists():
        return {}

    last: dict[str, Any] = {}
    for raw_line in path.read_text(encoding="utf-8").splitlines():
        line = raw_line.strip()
        if not line or line.startswith("#"):
            continue
        try:
            parsed = json.loads(line)
        except json.JSONDecodeError:
            continue
        if isinstance(parsed, dict):
            last = parsed
    return last


def deep_merge(base: dict[str, Any], override: dict[str, Any]) -> dict[str, Any]:
    merged = dict(base)
    for key, value in override.items():
        if isinstance(value, dict) and isinstance(merged.get(key), dict):
            merged[key] = deep_merge(merged[key], value)
        else:
            merged[key] = value
    return merged


def load_gaud_config(repo_root: Path) -> tuple[dict[str, Any], list[str]]:
    global_path = Path(os.environ.get("GAUD_CONFIG_GLOBAL", "~/.config/gaud.config.jsonl")).expanduser()
    repo_path = Path(os.environ.get("GAUD_CONFIG_REPO", str(repo_root / ".gaud.config.jsonl"))).expanduser()

    global_config = load_last_jsonl(global_path)
    repo_config = load_last_jsonl(repo_path)
    return deep_merge(global_config, repo_config), [str(global_path), str(repo_path)]


def normalize_provider(value: Any) -> str:
    if not isinstance(value, str):
        return "unknown"
    return PROVIDER_ALIASES.get(value.strip().lower(), value.strip().lower()) or "unknown"


def configured_role_entries(config: dict[str, Any]) -> list[dict[str, Any]]:
    entries: list[dict[str, Any]] = []

    orchestrator = config.get("orchestrator")
    if isinstance(orchestrator, dict):
        entries.append({
            "role": "orchestrator",
            "name": orchestrator.get("name") or "Orchestrator",
            "cli": orchestrator.get("cli"),
            "model": orchestrator.get("model"),
        })

    implementers = config.get("implementers")
    if isinstance(implementers, list):
        for index, item in enumerate(implementers, start=1):
            if isinstance(item, dict):
                entries.append({
                    "role": "implementer",
                    "name": item.get("name") or f"Implementer{index}",
                    "cli": item.get("cli"),
                    "model": item.get("model"),
                })

    if not entries:
        entries.extend([
            {"role": "orchestrator", "name": "Current session", "cli": "current-session", "model": None},
            {"role": "tpm", "name": "TPM", "cli": "claude", "model": None},
            {"role": "investigator", "name": "Investigator", "cli": "claude", "model": None},
            {"role": "ux-ui", "name": "UX/UI", "cli": "gemini", "model": None},
            {"role": "implementer", "name": "Implementer", "cli": "codex", "model": None},
            {"role": "integrator", "name": "Integrator", "cli": "opencode", "model": None},
        ])

    return entries


def fallback_role_entries(config: dict[str, Any], configured_entries: list[dict[str, Any]]) -> list[dict[str, Any]]:
    fallbacks = config.get("fallbacks")
    if not isinstance(fallbacks, dict):
        return []

    entries: list[dict[str, Any]] = []
    configured_keys = {
        (str(entry.get("role")), normalize_provider(entry.get("cli")))
        for entry in configured_entries
    }

    for role_name, values in fallbacks.items():
        if isinstance(values, str):
            candidates = [values]
        elif isinstance(values, list):
            candidates = values
        else:
            continue

        singular_role = "implementer" if role_name == "implementers" else str(role_name)
        for index, cli in enumerate(candidates, start=1):
            provider_id = normalize_provider(cli)
            key = (singular_role, provider_id)
            if key in configured_keys:
                continue
            configured_keys.add(key)
            entries.append({
                "role": f"fallback:{singular_role}",
                "name": f"Fallback {singular_role} {index}",
                "cli": cli,
                "model": None,
                "fallback_for": singular_role,
            })

    return entries


def role_entries(config: dict[str, Any]) -> list[dict[str, Any]]:
    configured_entries = configured_role_entries(config)
    return configured_entries + fallback_role_entries(config, configured_entries)


def usage_cache_candidates(repo_root: Path) -> list[Path]:
    explicit_values = [
        os.environ.get("GAUD_USAGE_CACHE"),
        os.environ.get("B2V_USAGE_CACHE"),
    ]
    values = explicit_values if any(explicit_values) else [
        str(repo_root / "usage-cache.json"),
        str(Path.cwd() / "usage-cache.json"),
        str(Path.home() / "Library/Application Support/back2vibing/usage-cache.json"),
        str(Path.home() / ".config/back2vibing/usage-cache.json"),
    ]
    seen: set[str] = set()
    paths: list[Path] = []
    for value in values:
        if not value:
            continue
        path = Path(value).expanduser()
        key = str(path)
        if key not in seen:
            seen.add(key)
            paths.append(path)
    return paths


def usage_state_is_fresh(parsed: dict[str, Any]) -> bool:
    last_refresh = parsed.get("last_refresh_ms")
    if not isinstance(last_refresh, int):
        return True
    return int(time.time() * 1000) - last_refresh <= MAX_USAGE_CACHE_AGE_MS


def load_usage_state(repo_root: Path) -> tuple[dict[str, Any] | None, str | None, list[str], list[str]]:
    checked = usage_cache_candidates(repo_root)
    stale: list[str] = []
    for path in checked:
        if not path.exists():
            continue
        try:
            parsed = json.loads(path.read_text(encoding="utf-8"))
        except (OSError, json.JSONDecodeError):
            continue
        if isinstance(parsed, dict) and isinstance(parsed.get("providers"), list):
            if usage_state_is_fresh(parsed):
                return parsed, str(path), [str(p) for p in checked], stale
            stale.append(str(path))
    return None, None, [str(p) for p in checked], stale


def quota_label(quota: dict[str, Any]) -> str:
    qtype = quota.get("quota_type")
    if isinstance(qtype, dict):
        kind = qtype.get("type")
        value = qtype.get("value")
    elif isinstance(qtype, str):
        kind = qtype
        value = None
    else:
        kind = None
        value = None

    labels = {
        "Session": "Session Limit",
        "session": "Session Limit",
        "Weekly": "Weekly Limit",
        "weekly": "Weekly Limit",
        "FiveHour": "Rolling 5-Hour",
        "five_hour": "Rolling 5-Hour",
        "Five Hour": "Rolling 5-Hour",
    }
    if kind in labels:
        return labels[kind]
    if kind == "ModelSpecific" and isinstance(value, str):
        return value
    return str(quota.get("id") or kind or "quota")


def normalize_match_text(value: Any) -> str:
    if not isinstance(value, str):
        return ""
    return "".join(ch for ch in value.lower() if ch.isalnum())


def quota_type_kind(quota: dict[str, Any]) -> str | None:
    qtype = quota.get("quota_type")
    if isinstance(qtype, dict):
        kind = qtype.get("type")
        return kind if isinstance(kind, str) else None
    if isinstance(qtype, str):
        return qtype
    return None


def quota_matches_model(quota: dict[str, Any], model: Any) -> bool:
    model_key = normalize_match_text(model)
    if not model_key:
        return False
    candidates = [quota.get("id"), quota_label(quota)]
    qtype = quota.get("quota_type")
    if isinstance(qtype, dict):
        candidates.append(qtype.get("value"))
    return any(model_key in normalize_match_text(candidate) or normalize_match_text(candidate) in model_key for candidate in candidates)


def select_relevant_quotas(quotas_raw: list[dict[str, Any]], model: Any) -> tuple[list[dict[str, Any]], str | None]:
    if not quotas_raw:
        return quotas_raw, None
    if not isinstance(model, str) or not model.strip():
        return quotas_raw, None

    model_specific = [quota for quota in quotas_raw if quota_type_kind(quota) == "ModelSpecific"]
    matching_model = [quota for quota in model_specific if quota_matches_model(quota, model)]
    non_model = [quota for quota in quotas_raw if quota_type_kind(quota) != "ModelSpecific"]

    if matching_model:
        return non_model + matching_model, f"Matched configured model {model}."
    if non_model:
        return non_model, f"No model-specific quota matched {model}; using provider-wide quotas."
    return quotas_raw, f"No model-specific quota matched {model}."


def provider_lookup(usage_state: dict[str, Any] | None) -> dict[str, dict[str, Any]]:
    lookup: dict[str, dict[str, Any]] = {}
    if not usage_state:
        return lookup

    for provider in usage_state.get("providers", []):
        if not isinstance(provider, dict):
            continue
        provider_id = normalize_provider(provider.get("provider_id"))
        lookup[provider_id] = provider
    return lookup


def classify_provider(provider_id: str, provider: dict[str, Any] | None, model: Any = None) -> dict[str, Any]:
    if provider_id in {"current-session", "current_session"}:
        return {
            "provider_id": provider_id,
            "provider_name": "Current session",
            "status": "ready",
            "confidence": "high",
            "remaining_percent": None,
            "reset_at_ms": None,
            "reset_text": None,
            "reason": "This conductor session is already running.",
            "quotas": [],
        }

    if provider is None:
        return {
            "provider_id": provider_id,
            "provider_name": KNOWN_PROVIDERS.get(provider_id, provider_id),
            "status": "unknown",
            "confidence": "low",
            "remaining_percent": None,
            "reset_at_ms": None,
            "reset_text": None,
            "reason": "No usage snapshot found for this provider.",
            "quotas": [],
        }

    error = provider.get("error")
    name = provider.get("provider_name") or KNOWN_PROVIDERS.get(provider_id, provider_id)
    all_quotas_raw = [q for q in provider.get("quotas", []) if isinstance(q, dict)]
    quotas_raw, model_note = select_relevant_quotas(all_quotas_raw, model)
    quotas = []
    for quota in quotas_raw:
        remaining = quota.get("percent_remaining")
        if isinstance(remaining, (int, float)):
            remaining_value = max(0.0, min(100.0, float(remaining)))
        else:
            remaining_value = None
        quotas.append({
            "id": quota.get("id"),
            "label": quota_label(quota),
            "percent_remaining": remaining_value,
            "reset_at_ms": quota.get("resets_at_ms") if isinstance(quota.get("resets_at_ms"), int) else None,
            "reset_text": quota.get("reset_text") if isinstance(quota.get("reset_text"), str) else None,
            "is_model_specific": quota_type_kind(quota) == "ModelSpecific",
        })

    finite_quotas = [q for q in quotas if isinstance(q.get("percent_remaining"), float)]
    limiting_quota = min(finite_quotas, key=lambda q: float(q["percent_remaining"])) if finite_quotas else None
    remaining_percent = float(limiting_quota["percent_remaining"]) if limiting_quota else None
    reset_at_ms = limiting_quota.get("reset_at_ms") if limiting_quota else None
    reset_text = limiting_quota.get("reset_text") if limiting_quota else None

    if isinstance(error, str) and error.strip():
        lower_error = error.lower()
        if "auth" in lower_error or "log" in lower_error:
            status = "auth-blocked"
        elif "rate" in lower_error or "limit" in lower_error or "quota" in lower_error:
            status = "rate-limited"
        elif "not available" in lower_error or "not found" in lower_error:
            status = "unavailable"
        else:
            status = "unknown"
        return {
            "provider_id": provider_id,
            "provider_name": name,
            "status": status,
            "confidence": "medium",
            "remaining_percent": remaining_percent,
            "reset_at_ms": reset_at_ms,
            "reset_text": reset_text,
            "reason": f"{error} {model_note or ''}".strip(),
            "auth_type": provider.get("auth_type"),
            "quotas": quotas,
        }

    if remaining_percent is None:
        status = "unknown"
        confidence = "low"
        reason = "Provider snapshot has no quota percentages."
        if model_note:
            reason = f"{reason} {model_note}"
    elif remaining_percent <= 0:
        status = "quota-blocked"
        confidence = "high"
        reason = "The tightest relevant quota is depleted."
        if model_note:
            reason = f"{reason} {model_note}"
    elif remaining_percent < 5:
        status = "rate-limited"
        confidence = "high"
        reason = "The tightest relevant quota has less than 5% remaining."
        if model_note:
            reason = f"{reason} {model_note}"
    else:
        status = "ready"
        confidence = "high"
        reason = "Usage snapshot reports quota remaining."
        if model_note:
            reason = f"{reason} {model_note}"

    return {
        "provider_id": provider_id,
        "provider_name": name,
        "status": status,
        "confidence": confidence,
        "remaining_percent": remaining_percent,
        "reset_at_ms": reset_at_ms,
        "reset_text": reset_text,
        "reason": reason,
        "auth_type": provider.get("auth_type"),
        "quotas": quotas,
    }


def reset_delta_minutes(reset_at_ms: Any) -> int | None:
    if not isinstance(reset_at_ms, int):
        return None
    return max(0, int((reset_at_ms - int(time.time() * 1000)) / 60000))


def recommendation(health: dict[str, Any]) -> str:
    status = health["status"]
    percent = health.get("remaining_percent")
    delta = reset_delta_minutes(health.get("reset_at_ms"))

    if status != "ready":
        if status == "unknown":
            return "usable if needed, but do not prefer without a fallback"
        return "avoid for this run unless the user explicitly chooses it"

    if isinstance(delta, int) and delta <= 90 and isinstance(percent, float) and percent >= 20:
        return "prefer soon: quota resets shortly and has enough room to spend"
    if isinstance(percent, float) and percent >= 50:
        return "strong candidate: healthy remaining usage"
    if isinstance(percent, float) and percent < 20:
        return "conserve: usable but close to critical"
    return "candidate"


def rank_key(item: dict[str, Any]) -> tuple[int, int, int, int, float]:
    health = item["health"]
    status = health["status"]
    percent = health.get("remaining_percent")
    percent_value = float(percent) if isinstance(percent, (int, float)) else -1.0
    delta = reset_delta_minutes(health.get("reset_at_ms"))
    delta_value = delta if isinstance(delta, int) else 999999

    if status != "ready":
        return (STATUS_RANK.get(status, 9), 1, 1, delta_value, -percent_value)

    conserve_bucket = 1 if 0 <= percent_value < 20 else 0
    expiring_bucket = 0 if percent_value >= 20 and delta_value <= 90 else 1
    return (STATUS_RANK["ready"], conserve_bucket, expiring_bucket, delta_value, -percent_value)


def build_report(repo_root: Path) -> dict[str, Any]:
    config, config_paths = load_gaud_config(repo_root)
    usage_state, usage_path, usage_paths, stale_usage_paths = load_usage_state(repo_root)
    providers = provider_lookup(usage_state)

    agents = []
    for entry in role_entries(config):
        provider_id = normalize_provider(entry.get("cli"))
        health = classify_provider(provider_id, providers.get(provider_id), entry.get("model"))
        agents.append({
            **entry,
            "provider_id": provider_id,
            "health": health,
            "recommendation": recommendation(health),
        })

    ranked = sorted(agents, key=rank_key)

    return {
        "config_paths": config_paths,
        "usage_cache_path": usage_path,
        "usage_cache_candidates": usage_paths,
        "stale_usage_cache_paths": stale_usage_paths,
        "has_config": bool(config),
        "has_usage_snapshot": usage_state is not None,
        "agents": agents,
        "ranked_agents": ranked,
    }


def fmt_percent(value: Any) -> str:
    if isinstance(value, (int, float)):
        return f"{value:.0f}%"
    return "unknown"


def fmt_reset(health: dict[str, Any]) -> str:
    delta = reset_delta_minutes(health.get("reset_at_ms"))
    if isinstance(delta, int):
        hours, minutes = divmod(delta, 60)
        if hours:
            return f"{hours}h {minutes}m"
        return f"{minutes}m"
    if health.get("reset_text"):
        return str(health["reset_text"])
    return "unknown"


def print_human(report: dict[str, Any]) -> None:
    print("GAUD_AGENT_USAGE")
    print(f"config_loaded={str(report['has_config']).lower()}")
    print(f"usage_snapshot={report['usage_cache_path'] or 'none'}")
    if report.get("stale_usage_cache_paths"):
        print(f"stale_usage_snapshots={','.join(report['stale_usage_cache_paths'])}")
    print("")
    print("Configured agents:")
    for agent in report["agents"]:
        health = agent["health"]
        print(
            "- "
            f"{agent.get('name')} [{agent.get('role')}] "
            f"cli={agent.get('cli') or 'unknown'} model={agent.get('model') or 'default'} "
            f"status={health['status']} remaining={fmt_percent(health.get('remaining_percent'))} "
            f"reset={fmt_reset(health)}"
        )
        print(f"  recommendation: {agent['recommendation']}")
        print(f"  reason: {health['reason']}")
    print("")
    print("Suggested order:")
    for index, agent in enumerate(report["ranked_agents"], start=1):
        health = agent["health"]
        print(
            f"{index}. {agent.get('name')} ({agent.get('cli') or 'unknown'}) - "
            f"{health['status']}, {fmt_percent(health.get('remaining_percent'))} left, "
            f"resets {fmt_reset(health)}"
        )
    print("")
    print("Ask the user which agents to launch using this summary; default to the suggested ready agents unless the user overrides.")


def main() -> int:
    parser = argparse.ArgumentParser(description="Show usage-aware gaud agent recommendations.")
    parser.add_argument("--repo", default=os.getcwd(), help="Repository root containing optional .gaud.config.jsonl")
    parser.add_argument("--json", action="store_true", help="Print machine-readable JSON")
    args = parser.parse_args()

    report = build_report(Path(args.repo).resolve())
    if args.json:
        print(json.dumps(report, indent=2, sort_keys=True))
    else:
        print_human(report)
    return 0


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