#!/usr/bin/env python3
"""Waybar custom module for Codex account rate limits."""

from __future__ import annotations

import argparse
import html
import json
import os
import selectors
import subprocess
import sys
import time
from datetime import datetime
from pathlib import Path
from string import Formatter
from typing import Any

FIVE_HOUR_MINS = 300
WEEKLY_MINS = 10_080
DEFAULT_FORMAT = "5H {five_hour_used}% · W {weekly_used}% · {five_hour_reset_in}"
DEFAULT_TOOLTIP = (
    "Codex ({plan})\n"
    "5-hour: {five_hour_used}% used, resets {five_hour_reset_at} ({five_hour_reset_in})\n"
    "Weekly: {weekly_used}% used, resets {weekly_reset_at} ({weekly_reset_in})"
)
ALLOWED_FIELDS = {
    "five_hour_used",
    "weekly_used",
    "five_hour_reset_in",
    "weekly_reset_in",
    "five_hour_reset_at",
    "weekly_reset_at",
    "plan",
}


class UsageError(RuntimeError):
    pass


def cache_path() -> Path:
    root = Path(os.environ.get("XDG_CACHE_HOME", Path.home() / ".cache"))
    return root / "codex-waybar" / "rate-limits.json"


def parse_json_lines(output: str, response_id: int = 2) -> dict[str, Any]:
    for line in output.splitlines():
        try:
            message = json.loads(line)
        except json.JSONDecodeError:
            continue
        if message.get("id") == response_id:
            if "error" in message:
                raise UsageError(str(message["error"]))
            result = message.get("result")
            if isinstance(result, dict):
                return result
            raise UsageError("Codex returned an invalid rate-limit result")
    raise UsageError("Codex did not return rate-limit data")


def fetch_rate_limits(timeout: float, codex_command: str = "codex") -> dict[str, Any]:
    requests = [
        {
            "id": 1,
            "method": "initialize",
            "params": {
                "clientInfo": {"name": "codex-waybar", "version": "1.0.0"},
            },
        },
        {"method": "initialized"},
        {"id": 2, "method": "account/rateLimits/read"},
    ]
    payload = "\n".join(json.dumps(item, separators=(",", ":")) for item in requests) + "\n"
    process = None
    try:
        process = subprocess.Popen(
            [codex_command, "app-server"],
            stdin=subprocess.PIPE,
            stdout=subprocess.PIPE,
            stderr=subprocess.PIPE,
            text=True,
            bufsize=1,
        )
        assert process.stdin is not None
        assert process.stdout is not None
        process.stdin.write(payload)
        process.stdin.flush()

        deadline = time.monotonic() + timeout
        selector = selectors.DefaultSelector()
        selector.register(process.stdout, selectors.EVENT_READ)
        while time.monotonic() < deadline:
            events = selector.select(deadline - time.monotonic())
            if not events:
                break
            line = process.stdout.readline()
            if not line:
                break
            try:
                message = json.loads(line)
            except json.JSONDecodeError:
                continue
            if message.get("id") == 2:
                return parse_json_lines(line)
        raise UsageError("Codex app-server timed out or returned no rate-limit data")
    except FileNotFoundError as exc:
        raise UsageError(f"{codex_command!r} was not found") from exc
    finally:
        if process is not None:
            if process.stdin:
                try:
                    process.stdin.close()
                except OSError:
                    pass
            if process.poll() is None:
                process.terminate()
            try:
                process.wait(timeout=2)
            except subprocess.TimeoutExpired:
                process.kill()
                process.wait()


def select_snapshot(result: dict[str, Any]) -> dict[str, Any]:
    by_id = result.get("rateLimitsByLimitId")
    if isinstance(by_id, dict):
        codex = by_id.get("codex")
        if isinstance(codex, dict):
            return codex
        for value in by_id.values():
            if isinstance(value, dict):
                return value
    snapshot = result.get("rateLimits")
    if isinstance(snapshot, dict):
        return snapshot
    raise UsageError("No Codex rate-limit snapshot is available")


def select_windows(snapshot: dict[str, Any]) -> tuple[dict[str, Any], dict[str, Any]]:
    windows = [
        window
        for key in ("primary", "secondary")
        if isinstance((window := snapshot.get(key)), dict)
    ]
    by_duration = {window.get("windowDurationMins"): window for window in windows}
    five_hour = by_duration.get(FIVE_HOUR_MINS)
    weekly = by_duration.get(WEEKLY_MINS)
    if not isinstance(five_hour, dict) or not isinstance(weekly, dict):
        raise UsageError("The 5-hour or weekly Codex limit window is unavailable")
    return five_hour, weekly


def load_cache(path: Path) -> dict[str, Any] | None:
    try:
        value = json.loads(path.read_text())
    except (OSError, json.JSONDecodeError):
        return None
    return value if isinstance(value, dict) else None


def save_cache(path: Path, result: dict[str, Any]) -> None:
    try:
        path.parent.mkdir(parents=True, exist_ok=True)
        temporary = path.with_suffix(".tmp")
        temporary.write_text(json.dumps({"saved_at": int(time.time()), "result": result}))
        temporary.replace(path)
    except OSError:
        pass


def human_countdown(resets_at: Any, now: int) -> str:
    if not isinstance(resets_at, int):
        return "unknown"
    remaining = max(0, resets_at - now)
    days, remainder = divmod(remaining, 86_400)
    hours, remainder = divmod(remainder, 3_600)
    minutes = (remainder + 59) // 60
    if minutes == 60:
        hours += 1
        minutes = 0
    if hours == 24:
        days += 1
        hours = 0
    if days:
        return f"{days}d {hours}h"
    if hours:
        return f"{hours}h {minutes:02d}m"
    return f"{minutes}m"


def reset_time(resets_at: Any) -> str:
    if not isinstance(resets_at, int):
        return "unknown"
    return datetime.fromtimestamp(resets_at).astimezone().strftime("%Y-%m-%d %H:%M")


def validate_template(template: str) -> None:
    try:
        fields = {field for _, field, _, _ in Formatter().parse(template) if field}
    except ValueError as exc:
        raise UsageError(f"Invalid format template: {exc}") from exc
    unknown = fields - ALLOWED_FIELDS
    if unknown:
        raise UsageError(f"Unknown format field(s): {', '.join(sorted(unknown))}")


def format_values(snapshot: dict[str, Any], now: int) -> tuple[dict[str, Any], int]:
    five_hour, weekly = select_windows(snapshot)
    five_used = int(five_hour.get("usedPercent", 0))
    weekly_used = int(weekly.get("usedPercent", 0))
    values = {
        "five_hour_used": five_used,
        "weekly_used": weekly_used,
        "five_hour_reset_in": human_countdown(five_hour.get("resetsAt"), now),
        "weekly_reset_in": human_countdown(weekly.get("resetsAt"), now),
        "five_hour_reset_at": reset_time(five_hour.get("resetsAt")),
        "weekly_reset_at": reset_time(weekly.get("resetsAt")),
        "plan": snapshot.get("planType") or "unknown",
    }
    return values, max(five_used, weekly_used)


def output_json(
    result: dict[str, Any],
    text_format: str,
    tooltip_format: str,
    warning: int,
    critical: int,
    stale: bool = False,
    stale_reason: str | None = None,
    now: int | None = None,
) -> dict[str, Any]:
    snapshot = select_snapshot(result)
    values, percentage = format_values(snapshot, now or int(time.time()))
    classes = ["critical" if percentage >= critical else "warning" if percentage >= warning else "normal"]
    if stale:
        classes.append("stale")
    tooltip = tooltip_format.format_map(values)
    if stale:
        tooltip += f"\n\nCached data: {stale_reason or 'refresh failed'}"
    return {
        "text": html.escape(text_format.format_map(values)),
        "tooltip": html.escape(tooltip),
        "class": classes,
        "percentage": percentage,
    }


def unavailable_json(reason: str) -> dict[str, Any]:
    return {
        "text": "Codex ?",
        "tooltip": html.escape(f"Codex usage unavailable\n{reason}"),
        "class": ["unavailable"],
        "percentage": 0,
    }


def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--format", default=DEFAULT_FORMAT, help="Waybar text template")
    parser.add_argument("--tooltip-format", default=DEFAULT_TOOLTIP, help="Waybar tooltip template")
    parser.add_argument("--warning", type=int, default=70, help="Warning threshold percentage")
    parser.add_argument("--critical", type=int, default=90, help="Critical threshold percentage")
    parser.add_argument("--timeout", type=float, default=10, help="Codex app-server timeout in seconds")
    parser.add_argument("--codex-command", default="codex", help=argparse.SUPPRESS)
    parser.add_argument("--no-cache", action="store_true", help="Disable cached fallback data")
    return parser.parse_args(argv)


def main(argv: list[str] | None = None) -> int:
    args = parse_args(argv)
    try:
        validate_template(args.format)
        validate_template(args.tooltip_format)
        if not 0 <= args.warning <= args.critical <= 100:
            raise UsageError("thresholds must satisfy 0 <= warning <= critical <= 100")
    except UsageError as exc:
        print(json.dumps(unavailable_json(str(exc)), separators=(",", ":")))
        return 2

    path = cache_path()
    try:
        result = fetch_rate_limits(args.timeout, args.codex_command)
        save_cache(path, result)
        output = output_json(result, args.format, args.tooltip_format, args.warning, args.critical)
    except UsageError as exc:
        cached = None if args.no_cache else load_cache(path)
        cached_result = cached.get("result") if isinstance(cached, dict) else None
        try:
            if not isinstance(cached_result, dict):
                raise UsageError(str(exc))
            output = output_json(
                cached_result,
                args.format,
                args.tooltip_format,
                args.warning,
                args.critical,
                stale=True,
                stale_reason=str(exc),
            )
        except UsageError:
            output = unavailable_json(str(exc))

    print(json.dumps(output, separators=(",", ":")))
    return 0


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