from __future__ import annotations

import json
import os
import re
import sys
import time
from pathlib import Path
from urllib.error import HTTPError, URLError
from urllib.parse import quote, unquote, urlencode
from urllib.request import Request, urlopen

import mysql.connector


ROOT_DIR = Path(__file__).resolve().parents[1]
DEFAULT_ENV_PATH = ROOT_DIR / ".env"
DEFAULT_OVERRIDE_PATH = ROOT_DIR / "scripts" / "data" / "plant_name_overrides.json"
LATIN_REGEX = re.compile(r"^[A-Za-z][A-Za-z .'\-]+$")
SCIENTIFIC_NAME_REGEX = re.compile(r"^[A-Z][a-z]+(?: [a-z][a-z\-]+){1,3}$")
HIRAGANA_REGEX = re.compile(r"[\u3040-\u309F]")
KATAKANA_REGEX = re.compile(r"[\u30A0-\u30FF]")
HAN_REGEX = re.compile(r"[\u3400-\u9FFF]")
GOOGLE_TRANSLATE_CACHE: dict[tuple[str, str, str], str] = {}


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 looks_latin(value: str) -> bool:
    return bool(LATIN_REGEX.fullmatch(normalize(value)))


def looks_scientific_name(value: str) -> bool:
    return bool(SCIENTIFIC_NAME_REGEX.fullmatch(normalize(value)))


def has_hiragana(value: str) -> bool:
    return bool(HIRAGANA_REGEX.search(value or ""))


def has_katakana(value: str) -> bool:
    return bool(KATAKANA_REGEX.search(value or ""))


def has_han(value: str) -> bool:
    return bool(HAN_REGEX.search(value or ""))


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


def translate_text(text: str, target_lang: str, source_lang: str = "auto", attempt: int = 1) -> str:
    key = (text, target_lang, source_lang)
    if key in GOOGLE_TRANSLATE_CACHE:
        return GOOGLE_TRANSLATE_CACHE[key]

    query = urlencode(
        [
            ("client", "gtx"),
            ("sl", source_lang),
            ("tl", target_lang),
            ("dt", "t"),
            ("q", text),
        ]
    )
    url = "https://translate.googleapis.com/translate_a/single?" + query
    request = Request(url, headers={"User-Agent": "Mozilla/5.0", "Connection": "close"})
    try:
        with urlopen(request, timeout=60) as response:
            data = json.loads(response.read().decode("utf-8"))
    except (HTTPError, URLError, TimeoutError):
        if attempt >= 5:
            raise
        time.sleep(attempt * 1.5)
        return translate_text(text, target_lang, source_lang=source_lang, attempt=attempt + 1)

    translated = "".join(part[0] for part in data[0] if part and part[0]).strip()
    GOOGLE_TRANSLATE_CACHE[key] = translated
    time.sleep(0.05)
    return translated


def safe_translate_text(text: str, target_lang: str, source_lang: str = "auto") -> str:
    try:
        return normalize(translate_text(text, target_lang, source_lang=source_lang))
    except Exception:
        return ""


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:
    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 fetch_active_rows_by_ko() -> dict[str, dict]:
    conn = create_db_connection()
    cursor = conn.cursor(dictionary=True)
    try:
        cursor.execute(
            """
            SELECT id, type_name, type_name_ko, type_name_en, type_name_ja, type_name_zh
            FROM plant_reminder_presets
            WHERE is_active = 1
            """
        )
        rows = cursor.fetchall()
        return {normalize(row["type_name_ko"]): row for row in rows if normalize(row["type_name_ko"])}
    finally:
        cursor.close()
        conn.close()


def ensure_override_rows(
    overrides: dict[str, dict[str, str]],
    rows_by_ko: dict[str, dict],
) -> int:
    created = 0
    for ko_name, row in rows_by_ko.items():
        entry_key = normalize(row.get("type_name_en", "")) or normalize(row.get("type_name", ""))
        if not entry_key or entry_key in overrides:
            continue
        overrides[entry_key] = {
            "ko": ko_name,
        }
        created += 1
    return created


def fetch_wikidata_locales(scientific_names: list[str], chunk_size: int = 70) -> dict[str, dict[str, list[str]]]:
    results: dict[str, dict[str, list[str]]] = {}
    unique_names: list[str] = []
    seen: set[str] = 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 ?labelJa ?labelZh ?aliasJa ?aliasZh ?jaArticle ?zhArticle WHERE {{
  VALUES ?scientificName {{ {values} }}
  ?item wdt:P225 ?scientificName .
  OPTIONAL {{ ?item rdfs:label ?labelJa FILTER(LANG(?labelJa) = "ja") }}
  OPTIONAL {{ ?item rdfs:label ?labelZh FILTER(LANG(?labelZh) = "zh") }}
  OPTIONAL {{ ?item skos:altLabel ?aliasJa FILTER(LANG(?aliasJa) = "ja") }}
  OPTIONAL {{ ?item skos:altLabel ?aliasZh FILTER(LANG(?aliasZh) = "zh") }}
  OPTIONAL {{ ?jaArticle schema:about ?item ; schema:isPartOf <https://ja.wikipedia.org/> . }}
  OPTIONAL {{ ?zhArticle schema:about ?item ; schema:isPartOf <https://zh.wikipedia.org/> . }}
}}
"""
        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 = normalize(row["scientificName"]["value"])
            entry = results.setdefault(
                scientific_name,
                {"ja": [], "zh": []},
            )

            for key, locale_key in (("labelJa", "ja"), ("aliasJa", "ja"), ("labelZh", "zh"), ("aliasZh", "zh")):
                value = normalize(row.get(key, {}).get("value", ""))
                if value and value != scientific_name and value not in entry[locale_key]:
                    entry[locale_key].append(value)

            for key, locale_key in (("jaArticle", "ja"), ("zhArticle", "zh")):
                article_url = normalize(row.get(key, {}).get("value", ""))
                if not article_url:
                    continue
                title = article_url.split("/wiki/")[-1]
                title = normalize(unquote(title).replace("_", " "))
                if title and title != scientific_name and title not in entry[locale_key]:
                    entry[locale_key].append(title)
        time.sleep(0.2)

    return results


def scientific_token_count(scientific_name: str) -> int:
    return len([token for token in re.split(r"[\s\-]+", normalize(scientific_name)) if token])


def looks_common_english_name(value: str, scientific_name: str) -> bool:
    value = normalize(value)
    if not value or not looks_latin(value):
        return False
    if value == scientific_name:
        return False
    return not looks_scientific_name(value)


def is_good_google_ja(value: str, scientific_name: str) -> bool:
    value = normalize(value)
    if not value or looks_latin(value):
        return False
    if has_hiragana(value):
        return False
    token_count = scientific_token_count(scientific_name)
    segment_count = len([segment for segment in re.split(r"[・･\s]+", value) if segment])
    if token_count >= 2 and segment_count < 2 and len(value) < token_count * 4:
        return False
    return has_katakana(value) or has_han(value)


def choose_ja_fallback(scientific_name: str) -> str:
    space_candidate = safe_translate_text(scientific_name, "ja", source_lang="en")
    if is_good_google_ja(space_candidate, scientific_name):
        return space_candidate

    hyphen_candidate = safe_translate_text(scientific_name.replace(" ", "-"), "ja", source_lang="en")
    if is_good_google_ja(hyphen_candidate, scientific_name):
        return hyphen_candidate

    return hyphen_candidate or space_candidate


def choose_ja_localized_name(scientific_name: str, english_name: str) -> str:
    common_english = normalize(english_name)
    if looks_common_english_name(common_english, scientific_name):
        common_candidate = safe_translate_text(common_english, "ja", source_lang="en")
        if is_good_google_ja(common_candidate, scientific_name):
            return common_candidate

    return choose_ja_fallback(scientific_name)


def is_good_zh(value: str) -> bool:
    value = normalize(value)
    return bool(value and not looks_latin(value) and has_han(value))


def choose_zh_fallback(scientific_name: str, english_name: str) -> str:
    candidates: list[str] = []
    if looks_common_english_name(english_name, scientific_name):
        candidates.append(safe_translate_text(english_name, "zh-CN", source_lang="en"))
    candidates.append(safe_translate_text(scientific_name, "zh-CN", source_lang="en"))
    candidates.append(safe_translate_text(scientific_name.replace(" ", "-"), "zh-CN", source_lang="en"))

    for candidate in candidates:
        if is_good_zh(candidate):
            return candidate
    for candidate in candidates:
        if candidate and not looks_latin(candidate):
            return candidate
    return ""


def pick_existing_localized(value: str, locale: str) -> str:
    value = normalize(value)
    if not value or looks_latin(value):
        return ""
    if locale == "ja" and not (has_hiragana(value) or has_katakana(value) or has_han(value)):
        return ""
    if locale == "zh" and not has_han(value):
        return ""
    return value


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

    load_env(str(DEFAULT_ENV_PATH))
    overrides = load_overrides(DEFAULT_OVERRIDE_PATH)
    rows_by_ko = fetch_active_rows_by_ko()
    seeded_overrides = ensure_override_rows(overrides, rows_by_ko)
    wikidata_locales = fetch_wikidata_locales(list(overrides.keys()))

    updated_ja = 0
    updated_zh = 0
    processed = 0
    dirty = False
    sample: list[dict[str, str]] = []

    for scientific_name, payload in overrides.items():
        scientific_name = normalize(scientific_name)
        ko_name = normalize(payload.get("ko", ""))
        row = rows_by_ko.get(ko_name, {})
        current_en = normalize(row.get("type_name_en", scientific_name)) or scientific_name
        wikidata_entry = wikidata_locales.get(scientific_name, {})

        existing_ja = pick_existing_localized(payload.get("ja", ""), "ja")
        current_ja = pick_existing_localized(row.get("type_name_ja", ""), "ja")
        if existing_ja:
            target_ja = existing_ja
        elif wikidata_entry.get("ja"):
            target_ja = normalize(wikidata_entry["ja"][0])
        elif current_ja:
            target_ja = current_ja
        else:
            target_ja = choose_ja_localized_name(scientific_name, current_en)

        existing_zh = pick_existing_localized(payload.get("zh", ""), "zh")
        current_zh = pick_existing_localized(row.get("type_name_zh", ""), "zh")
        if existing_zh:
            target_zh = existing_zh
        elif wikidata_entry.get("zh"):
            target_zh = normalize(wikidata_entry["zh"][0])
        elif current_zh:
            target_zh = current_zh
        else:
            target_zh = choose_zh_fallback(scientific_name, current_en)

        if target_ja and target_ja != normalize(payload.get("ja", "")):
            payload["ja"] = target_ja
            payload["ja_source"] = (
                "override"
                if existing_ja
                else "current_db"
                if current_ja
                else "wikidata"
                if wikidata_entry.get("ja")
                else "translate_fallback"
            )
            updated_ja += 1
            dirty = True

        if target_zh and target_zh != normalize(payload.get("zh", "")):
            payload["zh"] = target_zh
            payload["zh_source"] = (
                "override"
                if existing_zh
                else "current_db"
                if current_zh
                else "wikidata"
                if wikidata_entry.get("zh")
                else "translate_fallback"
            )
            updated_zh += 1
            dirty = True

        if len(sample) < 40 and (target_ja or target_zh):
            sample.append(
                {
                    "scientific_name": scientific_name,
                    "ko": ko_name,
                    "ja": target_ja,
                    "zh": target_zh,
                }
            )

        processed += 1
        if dirty and processed % 50 == 0:
            save_overrides(DEFAULT_OVERRIDE_PATH, overrides)

    save_overrides(DEFAULT_OVERRIDE_PATH, overrides)
    print(
        json.dumps(
                {
                    "seeded_overrides": seeded_overrides,
                    "updated_ja": updated_ja,
                    "updated_zh": updated_zh,
                    "sample": sample,
                },
            ensure_ascii=False,
            indent=2,
        )
    )


if __name__ == "__main__":
    main()
