from __future__ import annotations

import argparse
import csv
import json
import re
import sys
import time
from collections import Counter
from pathlib import Path
from urllib.error import HTTPError, URLError
from urllib.parse import quote, urlencode
from urllib.request import Request, urlopen

from bs4 import BeautifulSoup


ROOT_DIR = Path(__file__).resolve().parents[1]
DEFAULT_SOURCE_CSV = ROOT_DIR / "scripts" / "data" / "latin_fallback_presets.csv"
DEFAULT_OVERRIDE_PATH = ROOT_DIR / "scripts" / "data" / "plant_name_overrides.json"
GBIF_MATCH_URL = "https://api.gbif.org/v1/species/match?verbose=true&kingdom=Plantae&name="
GBIF_SPECIES_URL = "https://api.gbif.org/v1/species/"
KPNI_URL = "https://www.nature.go.kr/kpni/stndasrch/dtl/selectNtnStndaPlantList2.do"
HEADERS = {
    "User-Agent": "Mozilla/5.0",
    "Accept-Language": "ko,en-US;q=0.9,en;q=0.8",
    "Content-Type": "application/x-www-form-urlencoded",
}
INFRASPECIFIC_MARKERS = {"subsp.", "subsp", "ssp.", "ssp", "var.", "var", "f.", "f", "forma"}
QUOTE_PREFIXES = ("'", '"', "‘", "’", "“", "”")
REJECTED_KPNI_BASE_OVERRIDES = {
    "Fouquieria splendens": {"캄파눌라타스플렌덴스관봉옥"},
    "Heracleum sphondylium": {"엘레간스스폰딜리움어수리"},
}


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Fill plant name overrides with Korean names from KPNI after resolving accepted names via GBIF."
    )
    parser.add_argument("--source-csv", default=str(DEFAULT_SOURCE_CSV))
    parser.add_argument("--out", default=str(DEFAULT_OVERRIDE_PATH))
    parser.add_argument("--delay", type=float, default=0.15)
    parser.add_argument("--limit", type=int, default=0)
    return parser.parse_args()


def normalize(value: str) -> str:
    return " ".join((value or "").strip().split())


def strip_authors(full_name: str) -> str:
    tokens = normalize(full_name).split()
    if len(tokens) < 2:
        return normalize(full_name)

    kept: list[str] = []
    for index, token in enumerate(tokens):
        if index < 2:
            kept.append(token)
            continue

        lowered = token.lower()
        if lowered in INFRASPECIFIC_MARKERS:
            kept.append(token)
            continue

        if token.startswith("(") or re.match(r"^[A-Z]", token) or token.endswith("."):
            break

        kept.append(token)

    return " ".join(kept)


def extract_korean_base_name(korean_name: str) -> str:
    korean_name = normalize(korean_name)
    korean_name = re.sub(r"\s+[\"'‘’“”].*$", "", korean_name)
    korean_name = re.sub(r"\s+\(.*\)$", "", korean_name)
    return normalize(korean_name)


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


def load_overrides(path: Path) -> dict[str, dict[str, str]]:
    if not path.exists():
        return {}
    return json.loads(path.read_text(encoding="utf-8"))


def save_overrides(path: Path, payload: dict[str, dict[str, str]]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    ordered = {name: payload[name] for name in sorted(payload)}
    path.write_text(json.dumps(ordered, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")


def should_skip_existing(scientific_name: str, existing: dict[str, str]) -> bool:
    if not existing:
        return False
    current_ko = normalize(existing.get("ko", ""))
    source = normalize(existing.get("source", ""))

    if not current_ko or current_ko == scientific_name:
        return False

    if not source:
        return True

    return False


def fetch_json(url: str, retries: int = 4) -> dict:
    last_error: Exception | None = None
    for attempt in range(retries):
        try:
            request = Request(url, headers={"User-Agent": "Mozilla/5.0", "Connection": "close"})
            with urlopen(request, timeout=90) as response:
                return json.loads(response.read().decode("utf-8"))
        except (HTTPError, URLError, TimeoutError, ConnectionResetError) as exc:
            last_error = exc
            time.sleep(1.0 + attempt)
    raise RuntimeError(f"Failed to fetch JSON: {url}") from last_error


def fetch_kpni_rows(scientific_name: str, retries: int = 5) -> list[dict[str, str]]:
    payload = urlencode({"plantInfoSearchWrd": scientific_name}).encode("utf-8")
    last_error: Exception | None = None
    for attempt in range(retries):
        try:
            request = Request(KPNI_URL, data=payload, headers=HEADERS)
            with urlopen(request, timeout=90) as response:
                html = response.read().decode("utf-8", "ignore")

            soup = BeautifulSoup(html, "html.parser")
            tbody = soup.find("tbody")
            if not tbody:
                return []

            rows: list[dict[str, str]] = []
            for tr in tbody.find_all("tr"):
                cells = [normalize(td.get_text(" ", strip=True)) for td in tr.find_all("td")]
                if len(cells) < 5:
                    continue
                rows.append(
                    {
                        "kind": cells[0],
                        "status": cells[1],
                        "scientific_name": cells[2],
                        "korean_name": cells[3],
                        "updated_at": cells[4],
                    }
                )
            return rows
        except (HTTPError, URLError, TimeoutError, ConnectionResetError) as exc:
            last_error = exc
            time.sleep(1.0 + attempt)
    raise RuntimeError(f"Failed to fetch KPNI rows: {scientific_name}") from last_error


def resolve_gbif_accepted_name(scientific_name: str) -> dict[str, str]:
    data = fetch_json(GBIF_MATCH_URL + quote(scientific_name))
    match_type = normalize(str(data.get("matchType", "")))
    status = normalize(str(data.get("status", "")))
    canonical_name = normalize(data.get("canonicalName", "")) or scientific_name
    accepted_canonical = canonical_name

    if status == "SYNONYM":
        accepted_canonical = normalize(data.get("species", ""))
        if not accepted_canonical:
            accepted_usage_key = data.get("acceptedUsageKey")
            if accepted_usage_key:
                accepted = fetch_json(f"{GBIF_SPECIES_URL}{accepted_usage_key}")
                accepted_canonical = normalize(accepted.get("canonicalName", "")) or normalize(
                    strip_authors(accepted.get("scientificName", ""))
                )

    return {
        "match_type": match_type,
        "status": status,
        "canonical_name": canonical_name,
        "accepted_canonical": accepted_canonical or scientific_name,
    }


def pick_exact_species_match(query: str, rows: list[dict[str, str]]) -> dict[str, str] | None:
    query = normalize(query).lower()
    candidates: list[dict[str, str]] = []
    for row in rows:
        if not row.get("korean_name"):
            continue
        stripped = strip_authors(row.get("scientific_name", "")).lower()
        if stripped != query:
            continue
        candidates.append(row)

    if not candidates:
        return None

    def score(row: dict[str, str]) -> tuple[int, int, int]:
        return (
            1 if row.get("status") == "정명" else 0,
            1 if row.get("kind") == "자생식물" else 0,
            -len(row.get("scientific_name", "")),
        )

    candidates.sort(key=score, reverse=True)
    return candidates[0]


def is_consensus_row(query: str, scientific_name: str) -> tuple[bool, str]:
    scientific_name = normalize(scientific_name)
    if not scientific_name.startswith(query + " "):
        return False, ""

    suffix = scientific_name[len(query):].strip()
    if not suffix:
        return False, ""

    if suffix.startswith(QUOTE_PREFIXES):
        return True, "cultivar"

    head = suffix.split()[0].lower()
    if head in INFRASPECIFIC_MARKERS:
        return True, "infraspecific"

    return False, ""


def pick_consensus_base_name(query: str, rows: list[dict[str, str]]) -> dict[str, str] | None:
    candidates: list[dict[str, str]] = []
    kinds: set[str] = set()
    for row in rows:
        ok, candidate_kind = is_consensus_row(query, row.get("scientific_name", ""))
        if not ok:
            continue
        korean_base = extract_korean_base_name(row.get("korean_name", ""))
        if not korean_base:
            continue
        entry = dict(row)
        entry["korean_base"] = korean_base
        entry["candidate_kind"] = candidate_kind
        candidates.append(entry)
        kinds.add(candidate_kind)

    if not candidates:
        return None

    base_counts = Counter(row["korean_base"] for row in candidates)
    if len(base_counts) != 1:
        return None

    korean_base, count = next(iter(base_counts.items()))
    return {
        "korean_name": korean_base,
        "mode": "cultivar_consensus" if kinds == {"cultivar"} else "infraspecific_consensus",
        "row_count": str(count),
        "candidate_kind": ",".join(sorted(kinds)),
        "sample_scientific_name": candidates[0]["scientific_name"],
        "source_kind": candidates[0]["kind"],
        "source_status": candidates[0]["status"],
        "updated_at": candidates[0]["updated_at"],
    }


def score_candidate(candidate: dict[str, str]) -> tuple[int, int]:
    mode = candidate.get("mode", "")
    row_count = int(candidate.get("row_count", "1") or "1")
    if mode == "exact_species":
        return (3, row_count)
    if mode == "cultivar_consensus":
        return (2, row_count)
    if mode == "infraspecific_consensus":
        return (1, row_count)
    return (0, row_count)


def is_rejected_candidate(query_name: str, candidate: dict[str, str]) -> bool:
    rejected_names = REJECTED_KPNI_BASE_OVERRIDES.get(query_name, set())
    return normalize(candidate.get("ko", "")) in rejected_names


def find_best_kpni_candidate(query_name: str, search_name: str) -> dict[str, str] | None:
    rows = fetch_kpni_rows(search_name)
    exact = pick_exact_species_match(search_name, rows)
    if exact:
        candidate = {
            "ko": normalize(exact["korean_name"]),
            "mode": "exact_species",
            "search_name": search_name,
            "source_kind": normalize(exact["kind"]),
            "source_status": normalize(exact["status"]),
            "source_scientific_name": normalize(exact["scientific_name"]),
            "source_updated_at": normalize(exact["updated_at"]),
            "row_count": "1",
        }
        if not is_rejected_candidate(query_name, candidate):
            return candidate

    consensus = pick_consensus_base_name(search_name, rows)
    if consensus:
        candidate = {
            "ko": normalize(consensus["korean_name"]),
            "mode": consensus["mode"],
            "search_name": search_name,
            "source_kind": normalize(consensus["source_kind"]),
            "source_status": normalize(consensus["source_status"]),
            "source_scientific_name": normalize(consensus["sample_scientific_name"]),
            "source_updated_at": normalize(consensus["updated_at"]),
            "row_count": consensus["row_count"],
        }
        if not is_rejected_candidate(query_name, candidate):
            return candidate

    return None


def main() -> None:
    try:
        sys.stdout.reconfigure(encoding="utf-8")
    except Exception:
        pass

    args = parse_args()
    source_path = Path(args.source_csv)
    out_path = Path(args.out)

    source_rows = load_source_rows(source_path)
    if args.limit > 0:
        source_rows = source_rows[: args.limit]

    overrides = load_overrides(out_path)
    updated = 0
    already_present = 0
    no_match = 0
    sample: list[dict[str, str]] = []

    for index, row in enumerate(source_rows, start=1):
        scientific_name = normalize(row.get("type_name", ""))
        if not scientific_name:
            continue

        existing = overrides.get(scientific_name, {})
        if should_skip_existing(scientific_name, existing):
            already_present += 1
            continue

        gbif = resolve_gbif_accepted_name(scientific_name)
        search_names = [scientific_name]
        accepted_canonical = normalize(gbif.get("accepted_canonical", ""))
        if accepted_canonical and accepted_canonical != scientific_name:
            search_names.append(accepted_canonical)

        best_candidate: dict[str, str] | None = None
        for search_name in search_names:
            candidate = find_best_kpni_candidate(scientific_name, search_name)
            if not candidate:
                continue
            if best_candidate is None or score_candidate(candidate) > score_candidate(best_candidate):
                best_candidate = candidate

        if not best_candidate:
            no_match += 1
        else:
            overrides[scientific_name] = {
                **existing,
                "ko": normalize(best_candidate["ko"]),
                "source": "KPNI_GBIF",
                "source_mode": normalize(best_candidate["mode"]),
                "source_query_name": scientific_name,
                "source_search_name": normalize(best_candidate["search_name"]),
                "source_kind": normalize(best_candidate["source_kind"]),
                "source_status": normalize(best_candidate["source_status"]),
                "source_scientific_name": normalize(best_candidate["source_scientific_name"]),
                "source_updated_at": normalize(best_candidate["source_updated_at"]),
                "source_row_count": normalize(best_candidate["row_count"]),
                "gbif_match_type": normalize(gbif.get("match_type", "")),
                "gbif_status": normalize(gbif.get("status", "")),
                "gbif_canonical_name": normalize(gbif.get("canonical_name", "")),
                "gbif_accepted_canonical": accepted_canonical,
            }
            updated += 1
            if len(sample) < 40:
                sample.append(
                    {
                        "scientific_name": scientific_name,
                        "ko": overrides[scientific_name]["ko"],
                        "source_mode": overrides[scientific_name]["source_mode"],
                        "source_search_name": overrides[scientific_name]["source_search_name"],
                    }
                )
            save_overrides(out_path, overrides)

        if index % 50 == 0:
            print(
                json.dumps(
                    {
                        "progress": index,
                        "updated": updated,
                        "already_present": already_present,
                        "no_match": no_match,
                    },
                    ensure_ascii=False,
                )
            )
        if args.delay > 0:
            time.sleep(args.delay)

    save_overrides(out_path, overrides)
    print(
        json.dumps(
            {
                "source_csv": str(source_path),
                "out": str(out_path),
                "processed": len(source_rows),
                "updated": updated,
                "already_present": already_present,
                "no_match": no_match,
                "sample": sample,
            },
            ensure_ascii=False,
            indent=2,
        )
    )


if __name__ == "__main__":
    main()
