from __future__ import annotations

import argparse
from dataclasses import dataclass
from pathlib import Path
import re
import sys
from typing import Iterable


REGIONS = ("EU", "GR", "JP", "JPIO", "US", "FL")


@dataclass(frozen=True)
class FgiHit:
    region: str
    catalog: str
    fgi_file: Path
    oe: str
    reference: str
    figure_group: str


@dataclass(frozen=True)
class BziRecord:
    figure_group: str
    sequence: str
    image_code: str
    start_ym: str
    end_ym: str
    count_code: str
    option_code: str
    note_index: int


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


def format_oe(value: str) -> str:
    return f"{value[:5]}-{value[5:]}"


def iter_fgi_hits(epc_root: Path, oe: str) -> list[FgiHit]:
    hits: list[FgiHit] = []
    target = normalize_oe(oe)

    for region in REGIONS:
        region_dir = epc_root / region
        if not region_dir.is_dir():
            continue
        for catalog_dir in sorted(p for p in region_dir.iterdir() if p.is_dir()):
            for fgi_file in sorted(catalog_dir.glob("FGI*.DAT")):
                data = fgi_file.read_bytes()
                # FGI records are fixed-width 22-byte ASCII rows:
                # 0:10 OE without hyphen, 12:18 reference number, 18:22 figure group.
                for offset in range(0, len(data) - 21, 22):
                    row = data[offset : offset + 22].decode("ascii", errors="ignore")
                    if row[:10].strip() == target:
                        hits.append(
                            FgiHit(
                                region=region,
                                catalog=catalog_dir.name,
                                fgi_file=fgi_file,
                                oe=format_oe(target),
                                reference=row[12:18].strip(),
                                figure_group=row[18:22].strip(),
                            )
                        )
    return hits


def read_fgi_rows(fgi_file: Path) -> list[tuple[str, str, str]]:
    data = fgi_file.read_bytes()
    rows: list[tuple[str, str, str]] = []
    for offset in range(0, len(data) - 21, 22):
        row = data[offset : offset + 22].decode("ascii", errors="ignore")
        part = row[:10].strip()
        if not part:
            continue
        rows.append((format_oe(part), row[12:18].strip(), row[18:22].strip()))
    return rows


def same_position_candidates(hit: FgiHit) -> list[str]:
    candidates = {
        part
        for part, reference, figure_group in read_fgi_rows(hit.fgi_file)
        if reference == hit.reference and figure_group == hit.figure_group
    }
    return sorted(candidates)


def part_name(epc_root: Path, region: str, reference: str) -> str:
    region_dir = epc_root / region
    # HINMEI01 in this EPC package is Simplified Chinese (CP936/GBK).
    # Other HINMEI files in the sampled GR/US data are Western-language Latin text.
    language_files = (
        ("HINMEI01.DAT", "gbk"),
        ("HINMEI03.DAT", "latin1"),
        ("HINMEI04.DAT", "latin1"),
        ("HINMEI05.DAT", "latin1"),
    )
    for name, encoding in language_files:
        path = region_dir / name
        if not path.is_file():
            continue
        data = path.read_bytes()
        if len(data) % 66 != 0:
            continue
        for offset in range(0, len(data), 66):
            row = data[offset : offset + 66].decode(encoding, errors="ignore")
            if row[:6].strip() == reference:
                return row[6:].strip()
    return ""


def catalog_dir_for_hit(hit: FgiHit) -> Path:
    return hit.fgi_file.parent


def parse_bzi(catalog_dir: Path, figure_group: str) -> list[BziRecord]:
    records: list[BziRecord] = []
    for path in catalog_dir.glob("BZI*.DAT"):
        data = path.read_bytes()
        if len(data) % 37 != 0:
            continue
        for offset in range(0, len(data), 37):
            row = data[offset : offset + 37]
            text = row[:34].decode("ascii", errors="ignore")
            if text[:4] != figure_group:
                continue
            records.append(
                BziRecord(
                    figure_group=text[:4],
                    sequence=text[4:8],
                    image_code=text[8:15].strip(),
                    start_ym=text[15:21],
                    end_ym=text[21:27],
                    count_code=text[27:30],
                    option_code=text[30:34].strip(),
                    note_index=int.from_bytes(row[34:37], "little"),
                )
            )
    return records


def parse_emi(catalog_dir: Path) -> dict[str, int]:
    image_to_index: dict[str, int] = {}
    for path in catalog_dir.glob("EMI*.DAT"):
        data = path.read_bytes()
        count = len(data) // 12
        for index in range(count):
            row = data[index * 12 : (index + 1) * 12].decode("ascii", errors="ignore")
            image_code = row[-7:].strip()
            if image_code:
                image_to_index[image_code] = index
    return image_to_index


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


def extract_figure_images(hit: FgiHit, output_dir: Path) -> list[Path]:
    catalog_dir = catalog_dir_for_hit(hit)
    bzi_records = parse_bzi(catalog_dir, hit.figure_group)
    image_to_index = parse_emi(catalog_dir)
    emk_files = list(catalog_dir.glob("EMK*.DAT"))
    if not emk_files:
        return []
    emk_file = emk_files[0]
    data = emk_file.read_bytes()
    ranges = png_ranges(emk_file)
    output_dir.mkdir(parents=True, exist_ok=True)

    written: list[Path] = []
    for record in bzi_records:
        if record.image_code not in image_to_index:
            continue
        index = image_to_index[record.image_code]
        if index >= len(ranges):
            continue
        start, end = ranges[index]
        dst = (
            output_dir
            / f"{hit.region}_{hit.catalog}_{hit.figure_group}_{record.sequence}_{record.image_code}.png"
        )
        dst.write_bytes(data[start:end])
        written.append(dst)
    return written


def catalog_conditions(catalog_dir: Path) -> tuple[list[str], list[str]]:
    tkm_texts: list[str] = []
    kig_texts: list[str] = []
    for suffix in ("D03", "D01", "D04", "D05"):
        for path in catalog_dir.glob(f"TKM*.{suffix}"):
            text = path.read_bytes().decode("latin1", errors="ignore")
            if "POSITION" in text or "LHD" in text or "RHD" in text:
                tkm_texts.append(" ".join(text.split()))
        for path in catalog_dir.glob(f"KIG*.{suffix}"):
            text = path.read_bytes().decode("latin1", errors="ignore")
            for code in ("07LHD", "07RHD"):
                idx = text.find(code)
                if idx >= 0:
                    snippet = text[idx : idx + 80]
                    kig_texts.append(" ".join(snippet.split()))
    return tkm_texts[:3], sorted(set(kig_texts))


def catalog_frame_codes(catalog_dir: Path) -> list[str]:
    codes: set[str] = set()
    for path in catalog_dir.glob("KRK*.DAT"):
        text = path.read_bytes().decode("latin1", errors="ignore")
        for match in re.finditer(r"\b[A-Z]{2,4}\d{2,3}\b", text):
            codes.add(match.group(0))
    return sorted(codes)


def summarize_hits(hits: list[FgiHit]) -> str:
    if not hits:
        return "No FGI hits found."

    lines = []
    lines.append(f"OE: {hits[0].oe}")
    lines.append(f"Total hits: {len(hits)}")
    lines.append("")
    lines.append("region\tcatalog\treference\tfigure_group\tfgi_file")
    for hit in hits:
        lines.append(
            "\t".join(
                [
                    hit.region,
                    hit.catalog,
                    hit.reference,
                    hit.figure_group,
                    str(hit.fgi_file),
                ]
            )
        )
    return "\n".join(lines)


def summarize_deep(epc_root: Path, hits: list[FgiHit], output_dir: Path | None = None) -> str:
    if not hits:
        return ""

    lines: list[str] = []
    lines.append("")
    lines.append("Deep summary:")
    seen_catalogs: set[tuple[str, str, str]] = set()
    for hit in hits:
        key = (hit.region, hit.catalog, hit.figure_group)
        if key in seen_catalogs:
            continue
        seen_catalogs.add(key)
        catalog_dir = catalog_dir_for_hit(hit)
        matching_hits = [h for h in hits if h.region == hit.region and h.catalog == hit.catalog]
        refs = sorted({h.reference for h in matching_hits})
        names = [f"{ref}: {part_name(epc_root, hit.region, ref)}" for ref in refs]
        candidates = sorted({part for h in matching_hits for part in same_position_candidates(h)})
        tkm, kig = catalog_conditions(catalog_dir)
        frames = catalog_frame_codes(catalog_dir)
        bzi = parse_bzi(catalog_dir, hit.figure_group)
        lines.append("")
        lines.append(f"[{hit.region}/{hit.catalog}] figure_group={hit.figure_group}")
        lines.append(f"names: {'; '.join(names)}")
        lines.append(f"candidate_same_position_oe: {', '.join(candidates[:80])}")
        if len(candidates) > 80:
            lines.append(f"candidate_same_position_oe_more: {len(candidates) - 80}")
        if frames:
            lines.append(f"frame_codes: {', '.join(frames[:80])}")
        if tkm:
            lines.append(f"condition_fields: {' | '.join(tkm)}")
        if kig:
            lines.append(f"steering_values: {' | '.join(kig)}")
        if bzi:
            bzi_text = [
                f"{r.sequence}:{r.image_code}:{r.start_ym}-{r.end_ym}:{r.option_code}"
                for r in bzi
            ]
            lines.append(f"bzi_images: {', '.join(bzi_text)}")
        if output_dir:
            written = extract_figure_images(hit, output_dir)
            if written:
                lines.append("image_files: " + ", ".join(str(path) for path in written))
    return "\n".join(lines)


def catalog_contexts(epc_root: Path, region: str, catalogs: Iterable[str]) -> list[str]:
    region_dir = epc_root / region
    if not region_dir.is_dir():
        return []

    needles = {catalog.upper(): catalog.upper()[:5] for catalog in catalogs}
    results: list[str] = []
    for path in sorted(region_dir.iterdir()):
        if not path.is_file():
            continue
        if path.suffix.upper() not in {".DAT", ".ACD", ".INI"}:
            continue
        if path.name.upper() in {"DVDCTLG.DAT"}:
            continue
        # Avoid image bundles and other very large payloads during index probing.
        if path.stat().st_size > 80_000_000:
            continue
        data = path.read_bytes()
        text = data.decode("latin1", errors="ignore")
        for catalog, short_code in needles.items():
            for needle in {catalog, short_code}:
                start = 0
                while True:
                    idx = text.find(needle, start)
                    if idx < 0:
                        break
                    left = max(0, idx - 90)
                    right = min(len(text), idx + len(needle) + 180)
                    snippet = text[left:right].replace("\x00", " ")
                    snippet = " ".join(snippet.split())
                    results.append(f"{region}\t{catalog}\t{path.name}\t{idx}\t{snippet}")
                    start = idx + len(needle)
                    if len(results) >= 200:
                        return results
    return results


def main() -> int:
    parser = argparse.ArgumentParser(description="Query local Toyota EPC FGI files by OE.")
    parser.add_argument("oe", help="OE number, with or without hyphen")
    parser.add_argument(
        "--epc-root",
        default=r"D:\TMCEPCW3\epcdata",
        help="Toyota EPC epcdata root directory",
    )
    parser.add_argument(
        "--catalog-contexts",
        action="store_true",
        help="Also print compact catalog-code contexts from regional index files",
    )
    parser.add_argument(
        "--deep",
        action="store_true",
        help="Print part names, same-position candidates, catalog conditions, and image info",
    )
    parser.add_argument(
        "--extract-images",
        action="store_true",
        help="Extract matched figure PNG images into ./toyota_epc_output",
    )
    args = parser.parse_args()

    epc_root = Path(args.epc_root)
    hits = iter_fgi_hits(epc_root, args.oe)
    print(summarize_hits(hits))
    if args.catalog_contexts and hits:
        print("")
        print("Catalog contexts:")
        grouped: dict[str, set[str]] = {}
        for hit in hits:
            grouped.setdefault(hit.region, set()).add(hit.catalog)
        for region, catalogs in sorted(grouped.items()):
            for line in catalog_contexts(epc_root, region, sorted(catalogs)):
                print(line)
    if args.deep and hits:
        output_dir = Path("toyota_epc_output") if args.extract_images else None
        print(summarize_deep(epc_root, hits, output_dir))
    return 0


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