from __future__ import annotations

from collections.abc import Iterable, Mapping
from datetime import UTC, datetime

from health_client import AccountSnapshot
from limit_warning import WINDOW_5H, WINDOW_7D

_WINDOW_USED_PCT = {
    WINDOW_5H: lambda snapshot: snapshot.primary_used_pct,
    WINDOW_7D: lambda snapshot: snapshot.secondary_used_pct,
}

_WINDOW_RESET_AT = {
    WINDOW_5H: lambda snapshot: snapshot.primary_reset_at,
    WINDOW_7D: lambda snapshot: snapshot.secondary_reset_at,
}


def capped_windows(snapshot: AccountSnapshot, caps: Mapping[str, int]) -> frozenset[str]:
    capped: set[str] = set()
    for window, cap_pct in caps.items():
        used_pct = _WINDOW_USED_PCT[window](snapshot)
        if used_pct is not None and (100 - used_pct) <= cap_pct:
            capped.add(window)
    return frozenset(capped)


def describe_caps(snapshot: AccountSnapshot, caps: Mapping[str, int]) -> str:
    breached = capped_windows(snapshot, caps)
    if not breached:
        return ""
    parts: list[str] = []
    for window in (WINDOW_5H, WINDOW_7D):
        if window not in breached:
            continue
        used_pct = _WINDOW_USED_PCT[window](snapshot)
        assert used_pct is not None
        remaining = 100 - used_pct
        parts.append(f"{window}={remaining}% left <= {caps[window]}%")
    return "capped(" + ", ".join(parts) + ")"


def cap_reset_at(snapshot: AccountSnapshot, windows: Iterable[str]) -> str | None:
    reset_times: list[float] = []
    for window in windows:
        reset_at = _WINDOW_RESET_AT[window](snapshot)
        if reset_at is not None:
            reset_times.append(reset_at)
    if not reset_times:
        return None
    dt = datetime.fromtimestamp(int(min(reset_times)), tz=UTC)
    return dt.strftime("%Y-%m-%dT%H:%M:%SZ")
