from __future__ import annotations

import json
import os
import time
from pathlib import Path
from urllib.parse import quote
from urllib.request import Request, urlopen

import mysql.connector


ROOT_DIR = Path(__file__).resolve().parents[1]
DEFAULT_ENV_PATH = ROOT_DIR / ".env"
DATA_DIR = ROOT_DIR / "scripts" / "data"
OVERRIDE_PATH = DATA_DIR / "plant_name_overrides.json"


def load_env(env_path: str) -> None:
    path = Path(env_path)
    if not path.exists():
        return
    for line in path.read_text(encoding="utf-8").splitlines():
        line = line.strip()
        if not line or line.startswith("#") or "=" not in line:
            continue
        key, value = line.split("=", 1)
        os.environ.setdefault(key.strip(), value.strip())


def create_db_connection():
    return mysql.connector.connect(
        host=os.environ["DB_HOST"],
        port=int(os.environ.get("DB_PORT", "3306")),
        user=os.environ["DB_USER"],
        password=os.environ["DB_PASSWORD"],
        database=os.environ["DB_NAME"],
        charset="utf8mb4",
    )


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


def fetch_json(url: str) -> dict:
    req = Request(url, headers={"User-Agent": "Mozilla/5.0", "Connection": "close"})
    with urlopen(req, timeout=120) as response:
        return json.loads(response.read().decode("utf-8"))


def pick_better_label_entry(current: dict[str, str] | None, candidate: dict[str, str]) -> dict[str, str]:
    if current is None:
        return candidate

    def score(entry: dict[str, str]) -> tuple[int, int, int, int]:
        non_empty = sum(1 for key in ("ko", "en", "ja", "zh") if entry.get(key))
        return (
            1 if entry.get("ko") else 0,
            1 if entry.get("ja") else 0,
            1 if entry.get("zh") else 0,
            non_empty,
        )

    if score(candidate) > score(current):
        return candidate
    return current


def fetch_wikidata_labels(scientific_names: list[str], chunk_size: int = 80) -> dict[str, dict[str, str]]:
    labels: dict[str, dict[str, str]] = {}
    unique_names = []
    seen = set()
    for name in scientific_names:
        name = normalize(name)
        if name and name not in seen:
            unique_names.append(name)
            seen.add(name)

    for start in range(0, len(unique_names), chunk_size):
        chunk = unique_names[start:start + chunk_size]
        values = " ".join(f'"{name.replace(chr(34), r"\\\"")}"' for name in chunk)
        query = f"""
SELECT ?scientificName ?labelKo ?labelEn ?labelJa ?labelZh WHERE {{
  VALUES ?scientificName {{ {values} }}
  ?item wdt:P225 ?scientificName .
  OPTIONAL {{ ?item rdfs:label ?labelKo FILTER(LANG(?labelKo) = "ko") }}
  OPTIONAL {{ ?item rdfs:label ?labelEn FILTER(LANG(?labelEn) = "en") }}
  OPTIONAL {{ ?item rdfs:label ?labelJa FILTER(LANG(?labelJa) = "ja") }}
  OPTIONAL {{ ?item rdfs:label ?labelZh FILTER(LANG(?labelZh) = "zh") }}
}}
"""
        url = "https://query.wikidata.org/sparql?format=json&query=" + quote(query)
        data = fetch_json(url)
        for row in data.get("results", {}).get("bindings", []):
            scientific_name = row["scientificName"]["value"]
            candidate = {
                "ko": normalize(row.get("labelKo", {}).get("value", "")),
                "en": normalize(row.get("labelEn", {}).get("value", "")),
                "ja": normalize(row.get("labelJa", {}).get("value", "")),
                "zh": normalize(row.get("labelZh", {}).get("value", "")),
            }
            labels[scientific_name] = pick_better_label_entry(labels.get(scientific_name), candidate)
        time.sleep(0.2)
    return labels


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


def load_generated_records() -> list[dict]:
    records = []
    for path in sorted(DATA_DIR.glob("plant_reminder_catalog_inat*.ndjson")):
        with path.open("r", encoding="utf-8") as fp:
            for line in fp:
                if not line.strip():
                    continue
                row = json.loads(line)
                scientific_name = normalize(row.get("source_scientific_name", ""))
                current_ko = normalize((row.get("translations") or {}).get("ko", {}).get("type_name", row.get("type_name", "")))
                if not scientific_name or not current_ko:
                    continue
                records.append({
                    "current_ko": current_ko,
                    "scientific_name": scientific_name,
                })
    unique = {}
    for row in records:
        unique[row["current_ko"]] = row
    return list(unique.values())


def main() -> None:
    load_env(str(DEFAULT_ENV_PATH))
    generated = load_generated_records()
    scientific_names = [row["scientific_name"] for row in generated]
    wikidata_labels = fetch_wikidata_labels(scientific_names)
    name_overrides = load_name_overrides()

    conn = create_db_connection()
    cursor = conn.cursor(dictionary=True)
    updated = 0
    skipped = 0
    sample = []

    try:
        cursor.execute("SELECT id, type_name FROM plant_reminder_presets")
        current_rows = cursor.fetchall()
        current_name_to_id = {row["type_name"]: row["id"] for row in current_rows}
        reserved_names = {row["type_name"] for row in current_rows}

        planned = []
        for row in generated:
            current_ko = row["current_ko"]
            record_id = current_name_to_id.get(current_ko)
            if not record_id:
                continue
            scientific_name = row["scientific_name"]
            labels = wikidata_labels.get(scientific_name, {})
            override = name_overrides.get(scientific_name, {})
            target_ko = normalize(override.get("ko") or labels.get("ko") or scientific_name)
            target_en = normalize(override.get("en") or labels.get("en") or scientific_name)
            target_ja = normalize(override.get("ja") or labels.get("ja") or scientific_name)
            target_zh = normalize(override.get("zh") or labels.get("zh") or scientific_name)
            planned.append({
                "id": record_id,
                "current_ko": current_ko,
                "target_ko": target_ko,
                "target_en": target_en,
                "target_ja": target_ja,
                "target_zh": target_zh,
                "scientific_name": scientific_name,
            })

        planned_name_counts = {}
        for row in planned:
            planned_name_counts[row["target_ko"]] = planned_name_counts.get(row["target_ko"], 0) + 1

        for row in planned:
            target_ko = row["target_ko"]
            if planned_name_counts.get(target_ko, 0) > 1:
                target_ko = row["scientific_name"]
            if target_ko != row["current_ko"] and target_ko in reserved_names and target_ko not in {r["current_ko"] for r in planned}:
                target_ko = row["scientific_name"]

            cursor.execute(
                """
                UPDATE plant_reminder_presets
                SET type_name = %s,
                    type_name_ko = %s,
                    type_name_en = %s,
                    type_name_ja = %s,
                    type_name_zh = %s,
                    updated_at = NOW()
                WHERE id = %s
                """,
                (target_ko, target_ko, row["target_en"], row["target_ja"], row["target_zh"], row["id"]),
            )
            conn.commit()

            if target_ko != row["current_ko"] or row["target_en"] != row["scientific_name"]:
                updated += 1
                if len(sample) < 10:
                    sample.append({
                        "before": row["current_ko"],
                        "after": target_ko,
                        "scientific_name": row["scientific_name"],
                        "en": row["target_en"],
                        "ja": row["target_ja"],
                        "zh": row["target_zh"],
                    })
            else:
                skipped += 1
            reserved_names.add(target_ko)

        print(json.dumps({"updated": updated, "skipped": skipped, "sample": sample}, ensure_ascii=False, indent=2))
    finally:
        cursor.close()
        conn.close()


if __name__ == "__main__":
    main()
