#!/usr/bin/env python3
"""Extract the Plenti Market Lab's official regional inputs from raw ONS XLSX files.

Every workbook is hash-pinned. The extractor records the exact source sheet and
cell/range behind each output so the database can retain observation-level lineage.
"""

from __future__ import annotations

import hashlib
import json
from datetime import date, datetime
from pathlib import Path
from typing import Any

from openpyxl import load_workbook


ROOT = Path(__file__).resolve().parent
OUTPUT_PATH = ROOT / "data" / "official_extracted.json"

REGIONS = {
    "E12000001": {"name": "North East", "household_sheet": "North_East"},
    "E12000002": {"name": "North West", "household_sheet": "North_West"},
    "E12000003": {"name": "Yorkshire and The Humber", "household_sheet": "Yorkshire_and_The_Humber"},
    "E12000004": {"name": "East Midlands", "household_sheet": "East_Midlands"},
    "E12000005": {"name": "West Midlands", "household_sheet": "West_Midlands"},
    "E12000006": {"name": "East of England", "household_sheet": "East", "hpi_name": "East"},
    "E12000007": {"name": "London", "household_sheet": "London"},
    "E12000008": {"name": "South East", "household_sheet": "South_East"},
    "E12000009": {"name": "South West", "household_sheet": "South_West"},
}

EXPECTED_KEYS = {
    "population",
    "age_18_24",
    "age_25_34",
    "young_share",
    "density",
    "median_age",
    "households",
    "one_person_share",
    "unrelated_adults",
    "unrelated_share",
    "rent",
    "rent_growth",
    "one_bed_rent",
    "two_bed_rent",
    "flat_rent",
    "house_price",
    "house_price_growth",
}


def sha256(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as source:
        for chunk in iter(lambda: source.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def require_number(value: Any, label: str) -> float:
    if isinstance(value, bool) or not isinstance(value, (int, float)):
        raise ValueError(f"Expected a number for {label}; found {value!r}")
    return float(value)


def find_code_row(worksheet, code: str, *, start: int = 1, end: int | None = None) -> int:
    for (cell,) in worksheet.iter_rows(
        min_row=start,
        max_row=end or worksheet.max_row,
        min_col=1,
        max_col=1,
    ):
        if cell.value == code:
            return cell.row
    raise ValueError(f"Could not find {code} in {worksheet.title!r}")


def lineage(file_id: str, sheet: str, cell_range: str, transformation: str) -> dict[str, str]:
    return {
        "file_id": file_id,
        "sheet_name": sheet,
        "cell_range": cell_range,
        "transformation": transformation,
    }


def verify_files(raw_files: list[dict]) -> tuple[dict[str, Path], list[dict]]:
    paths: dict[str, Path] = {}
    verified: list[dict] = []
    for item in raw_files:
        path = ROOT / item["path"]
        if not path.is_file():
            raise FileNotFoundError(f"Missing raw source file: {path}")
        actual = sha256(path)
        if actual != item["sha256"]:
            raise ValueError(
                f"Hash mismatch for {path.name}: expected {item['sha256']}, found {actual}"
            )
        paths[item["file_id"]] = path
        verified.append({**item, "verified_sha256": actual})
    return paths, verified


def extract_population(paths: dict[str, Path], regions: dict, cell_lineage: dict) -> None:
    file_id = "population_mid_2025_xlsx"
    workbook = load_workbook(paths[file_id], data_only=True, read_only=True, keep_links=False)
    try:
        persons = workbook["MYE2 - Persons"]
        density = workbook["MYE5"]
        median = workbook["MYE6"]
        if persons["D8"].value != "All ages" or persons["W8"].value != "18" or persons["AM8"].value != "34":
            raise ValueError("Population workbook age headers changed")
        if density["F8"].value != "2025 people per sq. km":
            raise ValueError("Population-density header changed")
        if median["D8"].value != "Mid-2025":
            raise ValueError("Median-age period header changed")

        for code in REGIONS:
            person_row = find_code_row(persons, code, start=9, end=365)
            density_row = find_code_row(density, code, start=9, end=365)
            median_row = find_code_row(median, code, start=9, end=365)

            population = int(require_number(persons.cell(person_row, 4).value, f"{code} population"))
            age_18_24 = int(sum(require_number(persons.cell(person_row, column).value, f"{code} age") for column in range(23, 30)))
            age_25_34 = int(sum(require_number(persons.cell(person_row, column).value, f"{code} age") for column in range(30, 40)))
            regions[code].update(
                {
                    "population": population,
                    "age_18_24": age_18_24,
                    "age_25_34": age_25_34,
                    "young_share": round(100 * (age_18_24 + age_25_34) / population, 2),
                    "density": require_number(density.cell(density_row, 6).value, f"{code} density"),
                    "median_age": require_number(median.cell(median_row, 4).value, f"{code} median age"),
                }
            )
            cell_lineage[code].update(
                {
                    "population": [lineage(file_id, persons.title, f"D{person_row}", "direct value")],
                    "age_18_24": [lineage(file_id, persons.title, f"W{person_row}:AC{person_row}", "sum single-year ages 18–24")],
                    "age_25_34": [lineage(file_id, persons.title, f"AD{person_row}:AM{person_row}", "sum single-year ages 25–34")],
                    "young_share": [lineage(file_id, persons.title, f"D{person_row},W{person_row}:AM{person_row}", "100 × ages 18–34 / all ages; round to 2 decimals")],
                    "density": [lineage(file_id, density.title, f"F{density_row}", "direct value")],
                    "median_age": [lineage(file_id, median.title, f"D{median_row}", "direct value")],
                }
            )
    finally:
        workbook.close()


def extract_households(paths: dict[str, Path], regions: dict, cell_lineage: dict) -> None:
    size_file = "household_size_2025_xlsx"
    type_file = "household_types_2025_xlsx"
    size_workbook = load_workbook(paths[size_file], data_only=True, read_only=True, keep_links=False)
    type_workbook = load_workbook(paths[type_file], data_only=True, read_only=True, keep_links=False)
    try:
        for code, config in REGIONS.items():
            sheet_name = config["household_sheet"]
            size_sheet = size_workbook[sheet_name]
            type_sheet = type_workbook[sheet_name]
            if size_sheet["B13"].value != "2025 Estimate" or type_sheet["B12"].value != "2025 Estimate":
                raise ValueError(f"2025 household header changed in {sheet_name}")

            size_values = [require_number(size_sheet.cell(row, 2).value, f"{code} household size") for row in range(14, 20)]
            households = int(sum(size_values))
            one_person = size_values[0]
            unrelated = int(require_number(type_sheet["B16"].value, f"{code} unrelated adults"))
            regions[code].update(
                {
                    "households": households,
                    "one_person_share": round(100 * one_person / households, 2),
                    "unrelated_adults": unrelated,
                    "unrelated_share": round(100 * unrelated / households, 2),
                }
            )
            cell_lineage[code].update(
                {
                    "households": [lineage(size_file, sheet_name, "B14:B19", "sum household-size categories; thousands")],
                    "one_person_share": [lineage(size_file, sheet_name, "B14,B14:B19", "100 × one-person households / constructed total; round to 2 decimals")],
                    "unrelated_adults": [lineage(type_file, sheet_name, "B16", "direct value; thousands")],
                    "unrelated_share": [
                        lineage(type_file, sheet_name, "B16", "numerator: two or more unrelated adults; thousands"),
                        lineage(size_file, sheet_name, "B14:B19", "denominator: sum household-size categories; round result to 2 decimals"),
                    ],
                }
            )
    finally:
        size_workbook.close()
        type_workbook.close()


def extract_rents(paths: dict[str, Path], regions: dict, cell_lineage: dict) -> None:
    file_id = "private_rents_july_2026_xlsx"
    workbook = load_workbook(paths[file_id], data_only=True, read_only=True, keep_links=False)
    try:
        sheet = workbook["Table 1"]
        expected_headers = {
            1: "Time period",
            2: "Area code",
            7: "Annual change",
            8: "Rental price",
            12: "Rental price one bed",
            16: "Rental price two bed",
            40: "Rental price flat maisonette",
        }
        for column, header in expected_headers.items():
            if sheet.cell(3, column).value != header:
                raise ValueError(f"Private-rent header changed at column {column}")

        target_period = date(2026, 7, 1)
        found: set[str] = set()
        for cells in sheet.iter_rows(min_row=4, min_col=1, max_col=40):
            period_value = cells[0].value
            code = cells[1].value
            period = period_value.date() if isinstance(period_value, datetime) else period_value
            if code not in REGIONS or period != target_period:
                continue
            found.add(code)
            column_map = {
                "rent_growth": 7,
                "rent": 8,
                "one_bed_rent": 12,
                "two_bed_rent": 16,
                "flat_rent": 40,
            }
            for key, column in column_map.items():
                cell = cells[column - 1]
                regions[code][key] = require_number(cell.value, f"{code} {key}")
                cell_lineage[code][key] = [lineage(file_id, sheet.title, cell.coordinate, "direct value for July 2026")]
            if len(found) == len(REGIONS):
                break
        missing = set(REGIONS) - found
        if missing:
            raise ValueError(f"Private-rent rows missing for: {', '.join(sorted(missing))}")
    finally:
        workbook.close()


def extract_house_prices(paths: dict[str, Path], regions: dict, cell_lineage: dict) -> None:
    file_id = "house_prices_june_2026_xlsx"
    workbook = load_workbook(paths[file_id], data_only=True, read_only=True, keep_links=False)
    try:
        price_sheet = workbook["2"]
        growth_sheet = workbook["3"]
        header_to_column = {
            str(price_sheet.cell(3, column).value).strip(): column
            for column in range(2, 17)
        }
        for code, config in REGIONS.items():
            hpi_name = config.get("hpi_name", config["name"])
            if hpi_name not in header_to_column:
                raise ValueError(f"House-price column missing for {hpi_name}")

        target_rows: dict[str, int] = {}
        for label, sheet in (("price", price_sheet), ("growth", growth_sheet)):
            for row in range(4, 190):
                value = str(sheet.cell(row, 1).value or "").strip()
                if value.startswith("Jun 2026"):
                    target_rows[label] = row
                    break
            if label not in target_rows:
                raise ValueError(f"June 2026 row missing from HPI sheet {sheet.title}")

        for code, config in REGIONS.items():
            column = header_to_column[config.get("hpi_name", config["name"])]
            price_cell = price_sheet.cell(target_rows["price"], column)
            growth_cell = growth_sheet.cell(target_rows["growth"], column)
            regions[code]["house_price"] = int(require_number(price_cell.value, f"{code} house price"))
            regions[code]["house_price_growth"] = require_number(growth_cell.value, f"{code} house-price growth")
            cell_lineage[code]["house_price"] = [lineage(file_id, price_sheet.title, price_cell.coordinate, "direct value for June 2026")]
            cell_lineage[code]["house_price_growth"] = [lineage(file_id, growth_sheet.title, growth_cell.coordinate, "direct annual change for June 2026")]
    finally:
        workbook.close()


def extract_official_data(raw_files: list[dict], *, write_output: bool = True) -> dict:
    paths, verified_files = verify_files(raw_files)
    regions = {code: {"code": code} for code in REGIONS}
    cell_lineage = {code: {} for code in REGIONS}

    extract_population(paths, regions, cell_lineage)
    extract_households(paths, regions, cell_lineage)
    extract_rents(paths, regions, cell_lineage)
    extract_house_prices(paths, regions, cell_lineage)

    for code, row in regions.items():
        keys = set(row) - {"code"}
        if keys != EXPECTED_KEYS:
            raise ValueError(f"{code} output keys differ: expected {EXPECTED_KEYS}, found {keys}")
        if set(cell_lineage[code]) != EXPECTED_KEYS:
            raise ValueError(f"{code} does not have complete cell lineage")
    if len(regions) * len(EXPECTED_KEYS) != 153:
        raise ValueError("Expected exactly 153 official observations")

    result = {
        "meta": {
            "region_count": len(regions),
            "indicator_count": len(EXPECTED_KEYS),
            "observation_count": len(regions) * len(EXPECTED_KEYS),
            "hashes_verified": True,
        },
        "files": verified_files,
        "regions": list(regions.values()),
        "lineage": cell_lineage,
    }
    if write_output:
        OUTPUT_PATH.parent.mkdir(parents=True, exist_ok=True)
        OUTPUT_PATH.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8")
    return result


def main() -> None:
    snapshot = json.loads((ROOT / "data" / "source_snapshot.json").read_text(encoding="utf-8"))
    result = extract_official_data(snapshot["raw_files"])
    print(
        f"Extracted {result['meta']['observation_count']} observations "
        f"with cell lineage from {len(result['files'])} verified workbooks"
    )
    print(f"Built {OUTPUT_PATH.relative_to(ROOT)}")


if __name__ == "__main__":
    main()
