"""Profile a CSV or Parquet extract and print the findings as Markdown.

Usage: python profile_table.py <file.csv|file.parquet> [--as-of YYYY-MM-DD]

Prints aggregates only, never raw rows. Frequent values are shown only for columns
with few distinct values, so identifiers such as patient or referral numbers stay out
of the output. Parquet files need pyarrow installed.
"""

import argparse
from datetime import date
from pathlib import Path

import pandas as pd

TOP_N = 5
LOW_CARDINALITY = 50


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("path", type=Path)
    parser.add_argument("--as-of", type=date.fromisoformat, default=date.today())
    return parser.parse_args()


def load(path: Path) -> pd.DataFrame:
    if path.suffix == ".csv":
        return pd.read_csv(path, dtype=str, keep_default_na=False, na_values=[""])
    if path.suffix == ".parquet":
        return pd.read_parquet(path)
    raise SystemExit(f"Unsupported file type {path.suffix!r}: use .csv or .parquet")


def is_date_column(name: str) -> bool:
    return "date" in name.lower() or name.lower().endswith("_ts")


def typed(series: pd.Series) -> pd.Series:
    """Parse date-named columns as dates and numeric-looking columns as numbers."""
    if is_date_column(str(series.name)):
        return pd.to_datetime(series, errors="coerce", format="ISO8601")
    numbers = pd.to_numeric(series, errors="coerce")
    if numbers.notna().sum() == series.notna().sum():
        return numbers
    return series


def kind(series: pd.Series) -> str:
    if pd.api.types.is_datetime64_any_dtype(series):
        return "date"
    if pd.api.types.is_numeric_dtype(series):
        return "number"
    return "text"


def value_range(series: pd.Series) -> tuple[str, str]:
    if kind(series) == "text" or series.notna().sum() == 0:
        return "", ""
    low, high = series.min(), series.max()
    if kind(series) == "date":
        return low.date().isoformat(), high.date().isoformat()
    return f"{low:g}", f"{high:g}"


def column_table(raw: pd.DataFrame, data: pd.DataFrame) -> list[str]:
    lines = ["| Column | Type | Missing | Distinct | Min | Max |", "|---|---|---|---|---|---|"]
    for name in data.columns:
        col = data[name]
        missing = raw[name].isna().mean()
        low, high = value_range(col)
        lines.append(
            f"| {name} | {kind(col)} | {missing:.1%} | {raw[name].nunique():,} | {low} | {high} |"
        )
    return lines


def grain_findings(raw: pd.DataFrame) -> list[str]:
    rows = len(raw)
    keys = [c for c in raw.columns if raw[c].notna().all() and raw[c].nunique() == rows]
    lines = [f"- Fully duplicated rows: {int(raw.duplicated().sum()):,}"]
    if keys:
        lines.append(f"- Candidate keys (unique, never missing): {', '.join(keys)}")
        return lines
    best = max(raw.columns, key=lambda c: raw[c].nunique())
    lines.append(
        f"- No single-column key. Most distinct: {best}, "
        f"{raw[best].nunique():,} values in {rows:,} rows"
    )
    return lines


def text_findings(raw: pd.DataFrame, data: pd.DataFrame) -> list[str]:
    lines: list[str] = []
    for name in data.columns:
        if kind(data[name]) != "text" or raw[name].nunique() > LOW_CARDINALITY:
            continue
        top = raw[name].value_counts().head(TOP_N)
        shown = ", ".join(f"{value!r} {count:,}" for value, count in top.items())
        lines.append(f"- {name}: {shown}")
        values = raw[name].dropna().unique()
        cleaned = {str(v).strip().lower() for v in values}
        if len(cleaned) < len(values):
            lines.append(f"  - {len(values)} spellings collapse to {len(cleaned)} after trim/lower")
    return lines


def date_findings(raw: pd.DataFrame, data: pd.DataFrame, as_of: date) -> list[str]:
    lines: list[str] = []
    cutoff = pd.Timestamp(as_of)
    for name in data.columns:
        if kind(data[name]) != "date":
            continue
        unparsed = int(raw[name].notna().sum() - data[name].notna().sum())
        future = int((data[name] > cutoff).sum())
        lines.append(f"- {name}: {unparsed:,} unparseable, {future:,} after {as_of.isoformat()}")
    return lines


def profile(path: Path, as_of: date) -> str:
    raw = load(path)
    data = raw.apply(typed)
    sections = [
        f"# Profile of {path.name}",
        f"{len(raw):,} rows, {len(raw.columns)} columns.",
        "## Columns",
        "\n".join(column_table(raw, data)),
        "## Grain",
        "\n".join(grain_findings(raw)),
        "## Low-cardinality text columns",
        "\n".join(text_findings(raw, data)) or "- None",
        "## Dates",
        "\n".join(date_findings(raw, data, as_of)) or "- No date columns",
    ]
    return "\n\n".join(sections)


def main() -> None:
    args = parse_args()
    print(profile(args.path, args.as_of))


if __name__ == "__main__":
    main()
