from __future__ import annotations

import csv
import json
import statistics
from collections import defaultdict
from pathlib import Path


ROOT = Path(__file__).resolve().parent
DATA_DIR = ROOT / "data"
OUTPUT_DIR = ROOT / "outputs"
VALID_GRADES = {"CNS一等", "CNS二等"}
ACTIVE_STATUS = {"available"}


def latest_snapshot() -> Path:
    snapshots = sorted(DATA_DIR.glob("observed-offers-*.csv"))
    if not snapshots:
        raise FileNotFoundError("找不到 data/observed-offers-YYYY-MM-DD.csv")
    return snapshots[-1]


def parse_number(value: str | None) -> float | None:
    if value is None:
        return None
    text = value.strip().replace(",", "")
    if not text:
        return None
    return float(text)


def package_group(unit_weight_kg: float) -> str:
    if unit_weight_kg <= 1.5:
        return "≤1.5kg"
    if unit_weight_kg <= 2:
        return "1.5–2kg"
    if unit_weight_kg <= 3:
        return "2–3kg"
    if unit_weight_kg <= 5:
        return "3–5kg"
    return "超過5kg"


def read_rows(path: Path) -> list[dict[str, str]]:
    with path.open("r", encoding="utf-8-sig", newline="") as handle:
        return list(csv.DictReader(handle))


def write_rows(path: Path, rows: list[dict], fields: list[str]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", encoding="utf-8-sig", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=fields)
        writer.writeheader()
        writer.writerows(rows)


def describe(values: list[float]) -> dict[str, float | int | None]:
    if not values:
        return {
            "n": 0,
            "mean": None,
            "median": None,
            "min": None,
            "max": None,
        }
    return {
        "n": len(values),
        "mean": round(statistics.fmean(values), 4),
        "median": round(statistics.median(values), 4),
        "min": round(min(values), 4),
        "max": round(max(values), 4),
    }


def normalize(rows: list[dict[str, str]]) -> list[dict]:
    output: list[dict] = []
    for row in rows:
        unit_weight = parse_number(row.get("unit_weight_kg"))
        pack_count_number = parse_number(row.get("pack_count"))
        pack_count = int(pack_count_number) if pack_count_number is not None else None
        regular_price = parse_number(row.get("regular_price_ntd"))
        promo_price = parse_number(row.get("promo_price_ntd"))
        conditional_price = parse_number(row.get("conditional_price_ntd"))
        total_weight = (
            round(unit_weight * pack_count, 4)
            if unit_weight and pack_count and unit_weight > 0 and pack_count > 0
            else None
        )
        public_price = promo_price if promo_price is not None else regular_price
        public_per_kg = (
            round(public_price / total_weight, 4)
            if public_price and total_weight and public_price > 0 and total_weight > 0
            else None
        )
        discount_rate = (
            round(1 - promo_price / regular_price, 6)
            if promo_price is not None
            and regular_price is not None
            and regular_price > 0
            else None
        )
        eligible = (
            row.get("grade_claim") in VALID_GRADES
            and row.get("stock_status") in ACTIVE_STATUS
            and public_price is not None
            and public_price > 0
            and total_weight is not None
            and total_weight > 0
        )
        quality_flags: list[str] = []
        if row.get("grade_evidence_location") == "package_text":
            quality_flags.append("package_text_manual_review")
        if row.get("package_form") in {"", "未揭露"}:
            quality_flags.append("package_form_not_disclosed")
        if conditional_price is not None:
            quality_flags.append("conditional_price_excluded")
        if row.get("stock_status") not in ACTIVE_STATUS:
            quality_flags.append("not_currently_available")
        if not row.get("shipping_note"):
            quality_flags.append("shipping_not_observed")
        output.append(
            {
                **row,
                "total_weight_kg": total_weight,
                "public_price_ntd": public_price,
                "public_price_per_kg": public_per_kg,
                "discount_rate": discount_rate,
                "analysis_eligible": "TRUE" if eligible else "FALSE",
                "package_size_group": package_group(unit_weight) if unit_weight else "",
                "quality_flags": "|".join(quality_flags),
            }
        )
    return output


def grade_summary(rows: list[dict]) -> list[dict]:
    summary: list[dict] = []
    eligible = [row for row in rows if row["analysis_eligible"] == "TRUE"]
    for grade in ["CNS一等", "CNS二等"]:
        group = [row for row in eligible if row["grade_claim"] == grade]
        stats = describe([float(row["public_price_per_kg"]) for row in group])
        discounts = [
            float(row["discount_rate"])
            for row in group
            if row["discount_rate"] is not None
        ]
        summary.append(
            {
                "grade_claim": grade,
                "offer_count": len(group),
                "unique_product_count": len({row["canonical_product_id"] for row in group}),
                "source_count": len({row["source_name"] for row in group}),
                "mean_price_per_kg": stats["mean"],
                "median_price_per_kg": stats["median"],
                "min_price_per_kg": stats["min"],
                "max_price_per_kg": stats["max"],
                "mean_discount_rate": (
                    round(statistics.fmean(discounts), 6) if discounts else None
                ),
            }
        )
    return summary


def unique_product_grade_summary(rows: list[dict]) -> list[dict]:
    by_product: dict[str, list[dict]] = defaultdict(list)
    for row in rows:
        if row["analysis_eligible"] == "TRUE":
            by_product[row["canonical_product_id"]].append(row)

    collapsed = []
    for product_id, offers in by_product.items():
        grades = {row["grade_claim"] for row in offers}
        if len(grades) != 1:
            continue
        collapsed.append(
            {
                "canonical_product_id": product_id,
                "grade_claim": next(iter(grades)),
                "offer_count": len(offers),
                "median_offer_price_per_kg": round(
                    statistics.median(
                        float(row["public_price_per_kg"]) for row in offers
                    ),
                    4,
                ),
            }
        )

    summary = []
    for grade in ["CNS一等", "CNS二等"]:
        group = [row for row in collapsed if row["grade_claim"] == grade]
        stats = describe([row["median_offer_price_per_kg"] for row in group])
        summary.append(
            {
                "grade_claim": grade,
                "unique_product_count": len(group),
                "mean_product_price_per_kg": stats["mean"],
                "median_product_price_per_kg": stats["median"],
                "min_product_price_per_kg": stats["min"],
                "max_product_price_per_kg": stats["max"],
            }
        )
    return summary


def package_summary(rows: list[dict]) -> list[dict]:
    eligible = [row for row in rows if row["analysis_eligible"] == "TRUE"]
    groups = ["≤1.5kg", "1.5–2kg", "2–3kg", "3–5kg", "超過5kg"]
    output = []
    for size_group in groups:
        for grade in ["CNS一等", "CNS二等"]:
            group = [
                row
                for row in eligible
                if row["package_size_group"] == size_group
                and row["grade_claim"] == grade
            ]
            stats = describe([float(row["public_price_per_kg"]) for row in group])
            output.append(
                {
                    "package_size_group": size_group,
                    "grade_claim": grade,
                    "offer_count": len(group),
                    "median_price_per_kg": stats["median"],
                    "mean_price_per_kg": stats["mean"],
                    "min_price_per_kg": stats["min"],
                    "max_price_per_kg": stats["max"],
                }
            )
    return output


def source_summary(rows: list[dict]) -> list[dict]:
    output = []
    for source in sorted({row["source_name"] for row in rows}):
        group = [row for row in rows if row["source_name"] == source]
        eligible = [row for row in group if row["analysis_eligible"] == "TRUE"]
        output.append(
            {
                "source_name": source,
                "observed_offer_count": len(group),
                "eligible_offer_count": len(eligible),
                "grade_1_offer_count": sum(
                    row["grade_claim"] == "CNS一等" for row in eligible
                ),
                "grade_2_offer_count": sum(
                    row["grade_claim"] == "CNS二等" for row in eligible
                ),
                "sold_or_out_of_stock_count": sum(
                    row["stock_status"] != "available" for row in group
                ),
            }
        )
    return output


def matched_product_offers(rows: list[dict]) -> list[dict]:
    by_product: dict[str, list[dict]] = defaultdict(list)
    for row in rows:
        if row["analysis_eligible"] == "TRUE":
            by_product[row["canonical_product_id"]].append(row)
    output = []
    for product_id, offers in sorted(by_product.items()):
        if len({row["source_name"] for row in offers}) < 2:
            continue
        prices = [float(row["public_price_per_kg"]) for row in offers]
        for row in sorted(offers, key=lambda item: item["source_name"]):
            output.append(
                {
                    "canonical_product_id": product_id,
                    "product_name": row["product_name"],
                    "grade_claim": row["grade_claim"],
                    "source_name": row["source_name"],
                    "total_weight_kg": row["total_weight_kg"],
                    "public_price_ntd": row["public_price_ntd"],
                    "public_price_per_kg": row["public_price_per_kg"],
                    "within_product_min_price_per_kg": round(min(prices), 4),
                    "within_product_max_price_per_kg": round(max(prices), 4),
                    "url": row["url"],
                }
            )
    return output


def benchmark_summary() -> list[dict]:
    rows = read_rows(DATA_DIR / "official-benchmark-2024.csv")
    output = []
    for row in rows:
        weight = parse_number(row["unit_weight_kg"])
        count = parse_number(row["pack_count"])
        price = parse_number(row["price_ntd"])
        total_weight = weight * count if weight and count else None
        output.append(
            {
                **row,
                "total_weight_kg": round(total_weight, 4) if total_weight else None,
                "price_per_kg": (
                    round(price / total_weight, 4)
                    if price and total_weight
                    else None
                ),
            }
        )
    return output


def main() -> None:
    snapshot = latest_snapshot()
    raw_rows = read_rows(snapshot)
    processed = normalize(raw_rows)
    grades = grade_summary(processed)
    unique_products = unique_product_grade_summary(processed)
    packages = package_summary(processed)
    sources = source_summary(processed)
    matched = matched_product_offers(processed)
    benchmark = benchmark_summary()

    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    write_rows(
        OUTPUT_DIR / "processed-offers.csv",
        processed,
        list(processed[0].keys()),
    )
    write_rows(
        OUTPUT_DIR / "summary-by-grade.csv",
        grades,
        list(grades[0].keys()),
    )
    write_rows(
        OUTPUT_DIR / "summary-by-grade-unique-product.csv",
        unique_products,
        list(unique_products[0].keys()),
    )
    write_rows(
        OUTPUT_DIR / "summary-by-package.csv",
        packages,
        list(packages[0].keys()),
    )
    write_rows(
        OUTPUT_DIR / "summary-by-source.csv",
        sources,
        list(sources[0].keys()),
    )
    write_rows(
        OUTPUT_DIR / "matched-cross-source-offers.csv",
        matched,
        list(matched[0].keys()) if matched else [
            "canonical_product_id",
            "product_name",
            "grade_claim",
            "source_name",
            "total_weight_kg",
            "public_price_ntd",
            "public_price_per_kg",
            "within_product_min_price_per_kg",
            "within_product_max_price_per_kg",
            "url",
        ],
    )
    write_rows(
        OUTPUT_DIR / "official-benchmark-2024-processed.csv",
        benchmark,
        list(benchmark[0].keys()),
    )

    eligible = [row for row in processed if row["analysis_eligible"] == "TRUE"]
    results = {
        "snapshot_file": snapshot.name,
        "observed_offer_count": len(processed),
        "eligible_offer_count": len(eligible),
        "unique_product_count": len(
            {row["canonical_product_id"] for row in eligible}
        ),
        "source_count": len({row["source_name"] for row in processed}),
        "grade_summary_offer_level": grades,
        "grade_summary_unique_product_level": unique_products,
        "cross_source_matched_product_count": len(
            {row["canonical_product_id"] for row in matched}
        ),
        "official_benchmark_row_count": len(benchmark),
        "interpretation_limit": (
            "首輪便利樣本只描述指定來源公開供給快照；未控制重量、有機、產地、"
            "米種、品牌與通路前，不得把一二等差異解讀為 CNS 等級溢價。"
        ),
    }
    (OUTPUT_DIR / "analysis-results.json").write_text(
        json.dumps(results, ensure_ascii=False, indent=2),
        encoding="utf-8",
    )
    print(json.dumps(results, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
