from __future__ import annotations

import argparse
import json
from pathlib import Path
import sqlite3
import sys
from typing import Any


EPC_ROOT = Path(r"D:\TMCEPCW3\epcdata")
DB_PATH = Path(__file__).resolve().parent / "toyota_epc_output" / "toyota_epc_validation.sqlite"


def normalize_oe(value: str) -> str:
    compact = "".join(ch for ch in value.upper() if ch.isalnum())
    if len(compact) != 10:
        raise ValueError(f"OE should normalize to 10 characters, got {compact!r}")
    return f"{compact[:5]}-{compact[5:]}"


def rows_to_dicts(rows) -> list[dict[str, Any]]:
    return [dict(row) for row in rows]


def first_part_names(conn: sqlite3.Connection, market: str, reference_code: str) -> list[dict[str, Any]]:
    rows = conn.execute(
        """
        SELECT language_file, language_code, reference_code, part_name,
               source_file, row_no, byte_offset, record_len
        FROM epc_part_name
        WHERE market = ? AND reference_code = ?
        ORDER BY
          CASE language_file
            WHEN 'HINMEI01.DAT' THEN 0
            WHEN 'HINMEI03.DAT' THEN 1
            WHEN 'HINMEI04.DAT' THEN 2
            WHEN 'HINMEI05.DAT' THEN 3
            ELSE 9
          END,
          row_no
        LIMIT 8
        """,
        (market, reference_code),
    ).fetchall()
    return rows_to_dicts(rows)


def vehicle_names(conn: sqlite3.Connection, market: str, catalog: str) -> list[dict[str, Any]]:
    if not table_exists(conn, "epc_vehicle_name"):
        return []
    rows = conn.execute(
        """
        SELECT DISTINCT vehicle_name_epc, model_family,
               source_file, row_no, byte_offset, record_len
        FROM epc_vehicle_name
        WHERE market = ? AND catalog = ?
        ORDER BY row_no
        LIMIT 10
        """,
        (market, catalog),
    ).fetchall()
    return rows_to_dicts(rows)


def catalog_model_scope(
    conn: sqlite3.Connection, market: str, catalog: str, sample_limit: int
) -> dict[str, Any]:
    if not table_exists(conn, "epc_vehicle_detail"):
        return {
            "distinct_model_count": 0,
            "frame_codes": [],
            "engines": [],
            "model_samples": [],
        }

    counts = conn.execute(
        """
        SELECT
          COUNT(DISTINCT model) AS distinct_model_count,
          COUNT(DISTINCT frame_code) AS distinct_frame_code_count,
          COUNT(DISTINCT engine_epc) AS distinct_engine_count
        FROM epc_vehicle_detail
        WHERE market = ? AND catalog = ?
        """,
        (market, catalog),
    ).fetchone()
    frame_rows = conn.execute(
        """
        SELECT DISTINCT frame_code
        FROM epc_vehicle_detail
        WHERE market = ? AND catalog = ? AND frame_code <> ''
        ORDER BY frame_code
        LIMIT 120
        """,
        (market, catalog),
    ).fetchall()
    engine_rows = conn.execute(
        """
        SELECT DISTINCT engine_epc
        FROM epc_vehicle_detail
        WHERE market = ? AND catalog = ? AND engine_epc <> ''
        ORDER BY engine_epc
        LIMIT 80
        """,
        (market, catalog),
    ).fetchall()
    sample_rows = conn.execute(
        """
        SELECT DISTINCT model, frame_code, engine_epc, production_start,
                        production_end, steering, destination
        FROM epc_vehicle_detail
        WHERE market = ? AND catalog = ?
        ORDER BY frame_code, model, production_start
        LIMIT ?
        """,
        (market, catalog, sample_limit),
    ).fetchall()
    return {
        "distinct_model_count": int(counts["distinct_model_count"] or 0),
        "distinct_frame_code_count": int(counts["distinct_frame_code_count"] or 0),
        "distinct_engine_count": int(counts["distinct_engine_count"] or 0),
        "frame_codes": [row["frame_code"] for row in frame_rows],
        "engines": [row["engine_epc"] for row in engine_rows],
        "model_samples": rows_to_dicts(sample_rows),
    }


def same_position_oe(
    conn: sqlite3.Connection,
    market: str,
    catalog: str,
    reference_code: str,
    figure_group: str,
) -> list[str]:
    rows = conn.execute(
        """
        SELECT DISTINCT oe_no
        FROM epc_part_index
        WHERE market = ?
          AND catalog = ?
          AND reference_code = ?
          AND figure_group = ?
        ORDER BY oe_no
        """,
        (market, catalog, reference_code, figure_group),
    ).fetchall()
    return [row["oe_no"] for row in rows]


def figure_images(
    conn: sqlite3.Connection, market: str, catalog: str, figure_group: str
) -> list[dict[str, Any]]:
    rows = conn.execute(
        """
        SELECT
          b.sequence,
          b.image_code,
          b.start_ym,
          b.end_ym,
          b.count_code,
          b.option_code,
          b.note_index,
          b.source_file AS bzi_file,
          b.row_no AS bzi_row_no,
          b.byte_offset AS bzi_byte_offset,
          b.record_len AS bzi_record_len,
          MIN(i.source_file) AS emi_file,
          MIN(i.row_no) AS emi_row_no,
          MIN(i.byte_offset) AS emi_byte_offset,
          MIN(i.record_len) AS emi_record_len,
          MIN(i.image_index) AS image_index
        FROM epc_figure_image b
        LEFT JOIN epc_image_index i
          ON i.market = b.market
         AND i.catalog = b.catalog
         AND i.figure_group = b.figure_group
         AND i.image_code = b.image_code
        WHERE b.market = ? AND b.catalog = ? AND b.figure_group = ?
        GROUP BY b.id
        ORDER BY b.sequence, b.image_code
        """,
        (market, catalog, figure_group),
    ).fetchall()
    return rows_to_dicts(rows)


def table_exists(conn: sqlite3.Connection, table: str) -> bool:
    row = conn.execute(
        "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?",
        (table,),
    ).fetchone()
    return row is not None


PNG_CACHE: dict[Path, list[tuple[int, int]]] = {}


def png_ranges(path: Path) -> list[tuple[int, int]]:
    if path in PNG_CACHE:
        return PNG_CACHE[path]
    data = path.read_bytes()
    signature = b"\x89PNG\r\n\x1a\n"
    starts: list[int] = []
    start = 0
    while True:
        index = data.find(signature, start)
        if index < 0:
            break
        starts.append(index)
        start = index + 1
    ranges = [
        (starts[index], starts[index + 1] if index + 1 < len(starts) else len(data))
        for index in range(len(starts))
    ]
    PNG_CACHE[path] = ranges
    return ranges


def extract_images_for_group(
    group: dict[str, Any],
    images: list[dict[str, Any]],
    output_dir: Path,
) -> None:
    catalog_dir = EPC_ROOT / group["market"] / group["catalog"]
    emk_files = sorted(catalog_dir.glob("EMK*.DAT"))
    if not emk_files:
        for image in images:
            image["extracted_png"] = ""
            image["extract_note"] = "EMK file not found"
        return
    emk_file = emk_files[0]
    ranges = png_ranges(emk_file)
    data = emk_file.read_bytes()
    output_dir.mkdir(parents=True, exist_ok=True)
    for image in images:
        index = image.get("image_index")
        if index is None:
            image["extracted_png"] = ""
            image["extract_note"] = "image index not found"
            continue
        index = int(index)
        if index >= len(ranges):
            image["extracted_png"] = ""
            image["extract_note"] = "image index outside EMK payload"
            continue
        start, end = ranges[index]
        dst = output_dir / (
            f"{group['market']}_{group['catalog']}_{group['figure_group']}_"
            f"{image['sequence']}_{image['image_code']}.png"
        )
        dst.write_bytes(data[start:end])
        image["extracted_png"] = str(dst.resolve())
        image["extract_note"] = ""


def query_part(
    conn: sqlite3.Connection,
    oe_input: str,
    sample_limit: int,
    extract_images: bool,
    output_dir: Path,
) -> dict[str, Any]:
    oe_no = normalize_oe(oe_input)
    hit_count = conn.execute(
        "SELECT COUNT(*) AS c FROM epc_part_index WHERE oe_no = ?",
        (oe_no,),
    ).fetchone()["c"]
    group_rows = conn.execute(
        """
        SELECT market, catalog, reference_code, figure_group,
               COUNT(*) AS raw_hit_rows,
               MIN(source_file) AS sample_fgi_file,
               MIN(row_no) AS sample_fgi_row_no,
               MIN(byte_offset) AS sample_fgi_byte_offset
        FROM epc_part_index
        WHERE oe_no = ?
        GROUP BY market, catalog, reference_code, figure_group
        ORDER BY market, catalog, reference_code, figure_group
        """,
        (oe_no,),
    ).fetchall()

    groups: list[dict[str, Any]] = []
    for row in group_rows:
        group = dict(row)
        group["part_names"] = first_part_names(conn, group["market"], group["reference_code"])
        group["vehicle_names"] = vehicle_names(conn, group["market"], group["catalog"])
        group["catalog_vehicle_scope"] = catalog_model_scope(
            conn, group["market"], group["catalog"], sample_limit
        )
        group["candidate_same_position_oe"] = same_position_oe(
            conn,
            group["market"],
            group["catalog"],
            group["reference_code"],
            group["figure_group"],
        )
        images = figure_images(conn, group["market"], group["catalog"], group["figure_group"])
        if extract_images:
            extract_images_for_group(group, images, output_dir)
        group["figure_images"] = images
        group["raw_position_rule"] = {
            "oe_table": "epc_part_index",
            "same_position_key": "market + catalog + reference_code + figure_group",
            "warning": (
                "This is a same-catalog same-position candidate list. "
                "Final fitment still needs condition-file decoding."
            ),
        }
        groups.append(group)

    return {
        "input_oe": oe_input,
        "normalized_oe": oe_no,
        "database": str(DB_PATH.resolve()),
        "source": "local Toyota EPC raw files only",
        "hit_rows": int(hit_count),
        "hit_groups": len(groups),
        "groups": groups,
    }


def main() -> int:
    parser = argparse.ArgumentParser(description="Query the local Toyota EPC SQLite DB by OE.")
    parser.add_argument("oe", help="OE number, with or without hyphen")
    parser.add_argument("--db", default=str(DB_PATH), help="SQLite database path")
    parser.add_argument("--sample-limit", type=int, default=12, help="Model samples per catalog")
    parser.add_argument(
        "--extract-images",
        action="store_true",
        help="Extract original figure PNGs into the output folder",
    )
    parser.add_argument(
        "--output-dir",
        default=str(Path(__file__).resolve().parent / "toyota_epc_output" / "part_images"),
        help="Image output directory when --extract-images is set",
    )
    args = parser.parse_args()

    db_path = Path(args.db)
    if not db_path.exists():
        raise SystemExit(f"Database not found: {db_path}")
    conn = sqlite3.connect(db_path)
    conn.row_factory = sqlite3.Row
    try:
        result = query_part(
            conn,
            args.oe,
            args.sample_limit,
            args.extract_images,
            Path(args.output_dir),
        )
        print(json.dumps(result, ensure_ascii=False, indent=2))
    finally:
        conn.close()
    return 0


if __name__ == "__main__":
    if hasattr(sys.stdout, "reconfigure"):
        sys.stdout.reconfigure(encoding="utf-8", errors="replace")
    raise SystemExit(main())
