"""Bounded retry policy helper for a caller-verified transient 429."""
from datetime import datetime, timezone
import math
import random
import re
import time


_DIGITS = re.compile(r"[0-9]+", re.ASCII)
_IMF_FIXDATE = re.compile(
    r"(?P<weekday>Mon|Tue|Wed|Thu|Fri|Sat|Sun), "
    r"(?P<day>[0-9]{2}) (?P<month>Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec) "
    r"(?P<year>[0-9]{4}) (?P<hour>[0-9]{2}):(?P<minute>[0-9]{2}):(?P<second>[0-9]{2}) GMT",
    re.ASCII,
)
_RFC850_DATE = re.compile(
    r"(?P<weekday>Monday|Tuesday|Wednesday|Thursday|Friday|Saturday|Sunday), "
    r"(?P<day>[0-9]{2})-(?P<month>Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec)-"
    r"(?P<year>[0-9]{2}) (?P<hour>[0-9]{2}):(?P<minute>[0-9]{2}):(?P<second>[0-9]{2}) GMT",
    re.ASCII,
)
_ASCTIME_DATE = re.compile(
    r"(?P<weekday>Mon|Tue|Wed|Thu|Fri|Sat|Sun) "
    r"(?P<month>Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec) "
    r"(?P<day> [1-9]|[12][0-9]|3[01]) "
    r"(?P<hour>[0-9]{2}):(?P<minute>[0-9]{2}):(?P<second>[0-9]{2}) "
    r"(?P<year>[0-9]{4})",
    re.ASCII,
)
_MONTHS = {name: index for index, name in enumerate(
    ("Jan", "Feb", "Mar", "Apr", "May", "Jun", "Jul", "Aug", "Sep", "Oct", "Nov", "Dec"), 1)}
def _parse_http_date(value, now):
    match = _IMF_FIXDATE.fullmatch(value)
    kind = "imf"
    if match is None:
        match = _RFC850_DATE.fullmatch(value)
        kind = "rfc850"
    if match is None:
        match = _ASCTIME_DATE.fullmatch(value)
        kind = "asctime"
    if match is None:
        return None

    fields = match.groupdict()
    year = int(fields["year"])
    if kind == "rfc850":
        current_century = now.year - (now.year % 100)
        year = current_century + year
        if year - now.year > 50:
            year -= 100

    try:
        parsed = datetime(year, _MONTHS[fields["month"]], int(fields["day"]),
                          int(fields["hour"]), int(fields["minute"]),
                          int(fields["second"]), tzinfo=timezone.utc)
    except ValueError:
        return None

    expected_weekday = fields["weekday"]
    actual_weekday = (parsed.strftime("%a") if kind != "rfc850"
                      else parsed.strftime("%A"))
    if actual_weekday != expected_weekday:
        return None
    return max(0.0, (parsed - now).total_seconds())


def retry_after_seconds(value, now=None):
    """Parse RFC 9110 delay-seconds or HTTP-date; return None if invalid."""
    if not isinstance(value, str):
        return None
    raw = value.strip(" \t")
    if not raw:
        return None
    if _DIGITS.fullmatch(raw):
        digits = raw.lstrip("0") or "0"
        # A valid but enormous integer is represented as infinity so policy
        # defers it as over-budget instead of falling back to a shorter wait.
        if len(digits) > 308:
            return math.inf
        try:
            return float(int(digits))
        except (ValueError, OverflowError):
            return math.inf

    current = now or datetime.now(timezone.utc)
    if current.tzinfo is None:
        current = current.replace(tzinfo=timezone.utc)
    else:
        current = current.astimezone(timezone.utc)
    return _parse_http_date(raw, current)


def bounded_retry(send_once, *, retryable_codes, repeat_is_safe,
                  max_attempts=4, max_elapsed=45.0, max_server_wait=30.0,
                  base_delay=0.5, cap_delay=8.0, jitter_seconds=0.5,
                  sleep=time.sleep, clock=time.monotonic, rng=random.uniform):
    """Return (result, attempts); max_attempts includes the first request.

    send_once(timeout_s=...) makes exactly one request and returns status, code,
    headers and body. It enforces the supplied remaining-time budget.
    """
    if not repeat_is_safe:
        raise ValueError("retry requires a documented safe-to-repeat operation")
    if max_attempts < 1:
        raise ValueError("max_attempts includes the first attempt and must be >= 1")
    started = clock()
    deadline = started + max_elapsed
    last = None
    for attempt in range(1, max_attempts + 1):
        remaining = deadline - clock()
        if remaining <= 0:
            return last, attempt - 1
        status, code, headers, body = send_once(timeout_s=remaining)
        last = (status, code, headers, body)
        if status != 429 or code not in retryable_codes or attempt == max_attempts:
            return last, attempt

        remaining = deadline - clock()
        header_value = next((v for k, v in headers.items()
                             if k.lower() == "retry-after"), None)
        server_wait = retry_after_seconds(header_value)
        if server_wait is not None:
            if server_wait > max_server_wait:
                return last, attempt
            delay = server_wait + rng(0.0, jitter_seconds)
        else:
            ceiling = min(cap_delay, base_delay * (2 ** (attempt - 1)))
            delay = rng(0.0, ceiling)

        if delay >= remaining:
            return last, attempt
        sleep(delay)

    raise AssertionError("unreachable")
