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/"
NATURE_URL = "https://www.nature.go.kr/ekbi/plant/smpl/selectPlantSmplGnrlSrch1.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 = ("'", '"', "‘", "’", "“", "”")


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Fill plant name overrides with Korean names from Nature specimen search 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.12)
    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 specimen_count(value: str) -> int:
    digits = re.sub(r"[^0-9]", "", value or "")
    return int(digits or "0")


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_nature_rows(scientific_name: str, retries: int = 5) -> list[dict[str, str]]:
    payload = urlencode({"searchCnd": "1", "searchWrd": scientific_name}).encode("utf-8")
    last_error: Exception | None = None
    for attempt in range(retries):
        try:
            request = Request(NATURE_URL, data=payload, headers=HEADERS)
            with urlopen(request, timeout=90) as response:
                html = response.read().decode("utf-8", "ignore")

            soup = BeautifulSoup(html, "html.parser")
            rows: list[dict[str, str]] = []
            for tr in soup.find_all("tr"):
                cells = [normalize(td.get_text(" ", strip=True)) for td in tr.find_all("td")]
                if len(cells) != 4:
                    continue
                rows.append(
                    {
                        "korean_name": cells[0],
                        "scientific_name": cells[1],
                        "korean_family_name": cells[2],
                        "specimen_count": cells[3],
                    }
                )
            return rows
        except (HTTPError, URLError, TimeoutError, ConnectionResetError) as exc:
            last_error = exc
            time.sleep(1.0 + attempt)
    raise RuntimeError(f"Failed to fetch Nature specimen 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 = [row for row in rows if strip_authors(row.get("scientific_name", "")).lower() == query and row.get("korean_name")]
    if not candidates:
        return None
    candidates.sort(key=lambda row: specimen_count(row.get("specimen_count", "")), reverse=True)
    return candidates[0]


def is_consensus_row(query: str, scientific_name: str) -> bool:
    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
    head = suffix.split()[0].lower()
    return head in INFRASPECIFIC_MARKERS


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

    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,
        "row_count": str(count),
        "source_scientific_name": candidates[0]["scientific_name"],
        "source_family_name": candidates[0]["korean_family_name"],
        "source_specimen_count": candidates[0]["specimen_count"],
    }


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


def find_best_nature_candidate(search_name: str) -> dict[str, str] | None:
    rows = fetch_nature_rows(search_name)
    exact = pick_exact_species_match(search_name, rows)
    if exact:
        return {
            "ko": normalize(exact["korean_name"]),
            "mode": "exact_species",
            "search_name": search_name,
            "row_count": "1",
            "source_scientific_name": normalize(exact["scientific_name"]),
            "source_family_name": normalize(exact["korean_family_name"]),
            "source_specimen_count": normalize(exact["specimen_count"]),
        }

    consensus = pick_consensus_base_name(search_name, rows)
    if consensus:
        return {
            "ko": normalize(consensus["korean_name"]),
            "mode": "consensus",
            "search_name": search_name,
            "row_count": consensus["row_count"],
            "source_scientific_name": normalize(consensus["source_scientific_name"]),
            "source_family_name": normalize(consensus["source_family_name"]),
            "source_specimen_count": normalize(consensus["source_specimen_count"]),
        }

    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_nature_candidate(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": "NATURE_SPECIMEN_GBIF",
                "source_mode": normalize(best_candidate["mode"]),
                "source_query_name": scientific_name,
                "source_search_name": normalize(best_candidate["search_name"]),
                "source_scientific_name": normalize(best_candidate["source_scientific_name"]),
                "source_family_name": normalize(best_candidate["source_family_name"]),
                "source_specimen_count": normalize(best_candidate["source_specimen_count"]),
                "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()
