"""Readers for the Human Mortality Database (and sister Human Fertility Database)
flat text files (US extract)."""

import numpy as np
import pandas as pd


def _read_hmd(path):
    """Read a standard HMD 1x1 file (2-line header, whitespace-delimited).

    Age column has values 0..109 and an open group '110+'; it is returned as an
    integer age (110 for the open group). Numeric cells may be '.' (missing).
    """
    df = pd.read_csv(path, skiprows=2, sep=r"\s+", na_values=".")
    df["Age"] = df["Age"].astype(str).str.replace("+", "", regex=False).astype(int)
    return df


def read_mx(path):
    """Death rates m_x. Returns dict sex -> DataFrame[Year, Age] wide table of m_x."""
    df = _read_hmd(path)
    out = {}
    for sex in ("Female", "Male", "Total"):
        out[sex.lower()] = df.pivot(index="Year", columns="Age", values=sex)
    return out


def read_deaths(path):
    """Death counts. Returns dict sex -> DataFrame[Year, Age]."""
    df = _read_hmd(path)
    return {sex.lower(): df.pivot(index="Year", columns="Age", values=sex)
            for sex in ("Female", "Male", "Total")}


def read_population(path):
    """Population size (HMD estimate, Jan 1). Returns dict sex -> DataFrame[Year, Age]."""
    df = _read_hmd(path)
    return {sex.lower(): df.pivot(index="Year", columns="Age", values=sex)
            for sex in ("Female", "Male", "Total")}


def read_births(path):
    """Births by year and sex of the baby. Returns DataFrame indexed by Year."""
    return pd.read_csv(path, skiprows=2, sep=r"\s+", na_values=".").set_index("Year")


def read_e0(path):
    """Published life expectancy at birth. Returns DataFrame indexed by Year."""
    df = pd.read_csv(path, skiprows=2, sep=r"\s+", na_values=".")
    return df.set_index("Year")


def read_ltper(path):
    """Published period life table (bltper/fltper/mltper). Returns DataFrame."""
    return pd.read_csv(path, skiprows=1, sep=r"\s+", na_values=".").assign(
        Age=lambda d: d["Age"].astype(str).str.replace("+", "", regex=False).astype(int))


# ---------------------------------------------------------------------------
# Human Fertility Database (HFD) readers
# ---------------------------------------------------------------------------

def _hfd_age(series):
    """Map HFD age labels to integer age: '12-' -> 12 (open low), '55+' -> 55."""
    return (series.astype(str)
            .str.replace("-", "", regex=False)
            .str.replace("+", "", regex=False)
            .astype(int))


def read_asfr_triangles(path):
    """HFD period ASFR by Lexis triangle (USAasfrTR.txt): Year, Age, Cohort, ASFR.

    Each age-year square is split into two triangles (two cohorts); each triangle
    rate is on ~half-year exposure, so the two rows for an age-year sum to ~2x the
    annual rate. Returns the raw long DataFrame with an integer `Age`.
    """
    df = pd.read_csv(path, skiprows=2, sep=r"\s+", na_values=".")
    df["Age"] = _hfd_age(df["Age"])
    return df


def asfr_by_age(path):
    """Annual period ASFR by single year of age from the triangle file.

    Averages the two Lexis triangles per age-year (x0.5 of their sum) to recover
    the conventional one-year age-specific fertility rate.
    Returns DataFrame indexed by Year, columns = Age.
    """
    df = read_asfr_triangles(path)
    wide = 0.5 * df.groupby(["Year", "Age"])["ASFR"].sum().unstack("Age")
    return wide


def tfr_from_triangles(path):
    """Period Total Fertility Rate derived from the triangle ASFR file.

    TFR = sum over all ages of the annual ASFR = 0.5 * sum of all triangle rates.
    Returns a Series indexed by Year.
    """
    df = read_asfr_triangles(path)
    return 0.5 * df.groupby("Year")["ASFR"].sum()


def read_tfr(path):
    """HFD published period TFR (USAtfrRR.txt): columns TFR, TFR40. Indexed by Year."""
    return pd.read_csv(path, skiprows=2, sep=r"\s+", na_values=".").set_index("Year")
