from __future__ import annotations

import argparse
import json
import os
import re
import sys
import time
from pathlib import Path
from urllib.error import HTTPError, URLError
from urllib.parse import 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"
WIKI_BANNED_TITLES = ("属", "科", "目", "一覧", "리스트", "レッド", "植物", "数据", "列表")
LATIN_ONLY_REGEX = re.compile(r"^[A-Za-z0-9 .,'()/_-]+$")
SCIENTIFIC_NAME_REGEX = re.compile(r"^[A-Z][a-z]+(?: [a-z][a-z\\-]+){1,3}$")


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


def fetch_active_rows() -> list[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
            ORDER BY id
            """
        )
        return cursor.fetchall()
    finally:
        cursor.close()
        conn.close()


def fetch_json(url: str, *, retry_sleep: float = 2.0, attempts: int = 4) -> dict:
    request = Request(url, headers={"User-Agent": "Mozilla/5.0", "Connection": "close"})
    for attempt in range(1, attempts + 1):
        try:
            with urlopen(request, timeout=60) as response:
                return json.loads(response.read().decode("utf-8"))
        except HTTPError as exc:
            if exc.code == 429 and attempt < attempts:
                time.sleep(retry_sleep * attempt)
                continue
            raise
        except URLError:
            if attempt < attempts:
                time.sleep(retry_sleep * attempt)
                continue
            raise
    raise RuntimeError("unreachable")


def wiki_api(lang: str, params: dict[str, str | int]) -> dict:
    url = f"https://{lang}.wikipedia.org/w/api.php?" + urlencode(params)
    data = fetch_json(url)
    time.sleep(0.9)
    return data


def fetch_extract(lang: str, title: str) -> tuple[str, str]:
    data = wiki_api(
        lang,
        {
            "action": "query",
            "format": "json",
            "prop": "extracts",
            "explaintext": 1,
            "exintro": 1,
            "redirects": 1,
            "titles": title,
        },
    )
    pages = data.get("query", {}).get("pages", {})
    if not pages:
        return "", ""
    page = next(iter(pages.values()))
    return normalize(page.get("title", "")), normalize(page.get("extract", ""))


def choose_wiki_title(lang: str, scientific_name: str) -> str:
    data = wiki_api(
        lang,
        {
            "action": "query",
            "format": "json",
            "list": "search",
            "srsearch": scientific_name,
            "srlimit": 5,
            "srprop": "snippet",
        },
    )

    for result in data.get("query", {}).get("search", []):
        title = normalize(result.get("title", ""))
        if not title or any(token in title for token in WIKI_BANNED_TITLES):
            continue
        if looks_latin_only(title):
            continue

        resolved_title, extract = fetch_extract(lang, title)
        if not resolved_title or any(token in resolved_title for token in WIKI_BANNED_TITLES):
            continue
        if looks_latin_only(resolved_title):
            continue
        if scientific_name.lower() not in extract.lower():
            continue

        text_window = extract[:240]
        if lang == "ja" and "学名" not in text_window:
            continue
        if lang == "zh" and "学名" not in text_window and "學名" not in text_window:
            continue

        return resolved_title

    return ""


def iter_candidates(rows: list[dict], overrides: dict[str, dict[str, str]]) -> list[dict]:
    candidates: list[dict] = []
    for row in rows:
        scientific_name = normalize(row.get("type_name_en") or row.get("type_name") or "")
        if not scientific_name or not SCIENTIFIC_NAME_REGEX.fullmatch(scientific_name):
            continue
        payload = overrides.get(scientific_name)
        if not payload:
            continue
        if payload.get("ja_source") == "translate_fallback" or payload.get("zh_source") == "translate_fallback":
            candidates.append({"scientific_name": scientific_name, "ko": normalize(row.get("type_name_ko", ""))})
    return candidates


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--limit", type=int, default=25)
    args = parser.parse_args()

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

    load_env(str(DEFAULT_ENV_PATH))
    overrides = json.loads(DEFAULT_OVERRIDE_PATH.read_text(encoding="utf-8"))
    rows = fetch_active_rows()
    candidates = iter_candidates(rows, overrides)[: max(0, args.limit)]

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

    for item in candidates:
        scientific_name = item["scientific_name"]
        payload = overrides[scientific_name]

        if payload.get("ja_source") == "translate_fallback":
            ja_title = choose_wiki_title("ja", scientific_name)
            if ja_title and ja_title != normalize(payload.get("ja", "")):
                payload["ja"] = ja_title
                payload["ja_source"] = "jawiki_search"
                updated_ja += 1

        if payload.get("zh_source") == "translate_fallback":
            zh_title = choose_wiki_title("zh", scientific_name)
            if zh_title and zh_title != normalize(payload.get("zh", "")):
                payload["zh"] = zh_title
                payload["zh_source"] = "zhwiki_search"
                updated_zh += 1

        if len(sample) < 50:
            sample.append(
                {
                    "scientific_name": scientific_name,
                    "ko": item["ko"],
                    "ja": normalize(payload.get("ja", "")),
                    "zh": normalize(payload.get("zh", "")),
                    "ja_source": normalize(payload.get("ja_source", "")),
                    "zh_source": normalize(payload.get("zh_source", "")),
                }
            )

    ordered = {name: overrides[name] for name in sorted(overrides)}
    DEFAULT_OVERRIDE_PATH.write_text(json.dumps(ordered, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")

    print(
        json.dumps(
            {
                "processed": len(candidates),
                "updated_ja": updated_ja,
                "updated_zh": updated_zh,
                "sample": sample,
            },
            ensure_ascii=False,
            indent=2,
        )
    )


if __name__ == "__main__":
    main()
