from __future__ import annotations

import argparse
import json
import os
import re
import sys
import time
from pathlib import Path
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"
LATIN_ONLY_REGEX = re.compile(r"^[A-Za-z0-9 .,'()/_-]+$")
KATAKANA_TRANSLITERATION_REGEX = re.compile(r"^[\u30A0-\u30FFー・\s]+$")
SCIENTIFIC_NAME_REGEX = re.compile(r"^[A-Z][a-z-]+(?: [a-z×x.-]+){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 looks_like_japanese_transliteration(value: str) -> bool:
    normalized = normalize(value)
    return bool(normalized) and "・" in normalized and bool(KATAKANA_TRANSLITERATION_REGEX.fullmatch(normalized))


def looks_like_scientific_name(value: str) -> bool:
    return bool(SCIENTIFIC_NAME_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) -> dict:
    request = Request(url, headers={"User-Agent": "Mozilla/5.0", "Connection": "close"})
    with urlopen(request, timeout=60) as response:
        return json.loads(response.read().decode("utf-8"))


def choose_inat_common_name(scientific_name: str, locale: str) -> str:
    params = urlencode(
        {
            "q": scientific_name,
            "locale": locale,
            "per_page": 10,
        }
    )
    url = "https://api.inaturalist.org/v1/taxa/autocomplete?" + params
    data = fetch_json(url)
    scientific_name_norm = normalize(scientific_name).lower()
    scientific_name_like = looks_like_scientific_name(scientific_name)

    best_score = -1
    best_value = ""
    for item in data.get("results", []):
        candidate_name = normalize(item.get("name", ""))
        matched_term = normalize(item.get("matched_term", ""))
        preferred_common_name = normalize(item.get("preferred_common_name", ""))
        if not preferred_common_name or looks_latin_only(preferred_common_name):
            continue
        candidate_name_norm = candidate_name.lower()
        matched_term_norm = matched_term.lower()
        score = -1
        if scientific_name_like:
            if candidate_name_norm == scientific_name_norm or matched_term_norm == scientific_name_norm:
                score = 3
        else:
            if candidate_name_norm == scientific_name_norm or matched_term_norm == scientific_name_norm:
                score = 3
            elif candidate_name_norm.startswith(scientific_name_norm + " ") or matched_term_norm.startswith(scientific_name_norm + " "):
                score = 2
            elif scientific_name_norm in candidate_name_norm or scientific_name_norm in matched_term_norm:
                score = 1

        if score < 0:
            continue

        if score > best_score or (score == best_score and len(preferred_common_name) < len(best_value or preferred_common_name + " ")):
            best_score = score
            best_value = preferred_common_name

    time.sleep(0.15)
    return best_value


def iter_candidate_rows(
    rows: list[dict],
    overrides: dict[str, dict[str, str]],
    offset: int,
    limit: int,
    include_ja_transliterations: bool,
    transliteration_sources: set[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:
            continue
        payload = overrides.get(scientific_name)
        if not payload:
            continue
        should_include_ja = payload.get("ja_source") == "translate_fallback"
        if (
            not should_include_ja
            and include_ja_transliterations
            and payload.get("ja_source") in transliteration_sources
            and looks_like_scientific_name(scientific_name)
            and looks_like_japanese_transliteration(payload.get("ja", ""))
        ):
            should_include_ja = True
        if should_include_ja or payload.get("zh_source") == "translate_fallback":
            candidates.append(row)

    if offset > 0:
        candidates = candidates[offset:]
    if limit > 0:
        candidates = candidates[:limit]
    return candidates


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--offset", type=int, default=0)
    parser.add_argument("--limit", type=int, default=0)
    parser.add_argument("--include-ja-transliterations", action="store_true")
    parser.add_argument("--ja-transliteration-sources", default="current_db,inaturalist,wikidata")
    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"))
    transliteration_sources = {
        normalize(source)
        for source in args.ja_transliteration_sources.split(",")
        if normalize(source)
    }
    rows = iter_candidate_rows(
        fetch_active_rows(),
        overrides,
        args.offset,
        args.limit,
        include_ja_transliterations=args.include_ja_transliterations,
        transliteration_sources=transliteration_sources,
    )

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

    for row in rows:
        scientific_name = normalize(row.get("type_name_en") or row.get("type_name") or "")
        if not scientific_name:
            continue

        payload = overrides.get(scientific_name)
        if not payload:
            continue

        changed = False

        should_refresh_ja = payload.get("ja_source") == "translate_fallback"
        if (
            not should_refresh_ja
            and args.include_ja_transliterations
            and payload.get("ja_source") in transliteration_sources
            and looks_like_scientific_name(scientific_name)
            and looks_like_japanese_transliteration(payload.get("ja", ""))
        ):
            should_refresh_ja = True

        if should_refresh_ja:
            ja_name = choose_inat_common_name(scientific_name, "ja")
            if ja_name and ja_name != normalize(payload.get("ja", "")):
                payload["ja"] = ja_name
                payload["ja_source"] = "inaturalist"
                updated_ja += 1
                changed = True

        if payload.get("zh_source") == "translate_fallback":
            zh_name = choose_inat_common_name(scientific_name, "zh-CN")
            if zh_name and zh_name != normalize(payload.get("zh", "")):
                payload["zh"] = zh_name
                payload["zh_source"] = "inaturalist"
                updated_zh += 1
                changed = True

        if changed and len(sample) < 80:
            sample.append(
                {
                    "scientific_name": scientific_name,
                    "ko": normalize(row.get("type_name_ko", "")),
                    "ja": normalize(payload.get("ja", "")),
                    "zh": normalize(payload.get("zh", "")),
                }
            )
        elif not changed:
            skipped += 1

    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(rows),
                "updated_ja": updated_ja,
                "updated_zh": updated_zh,
                "skipped": skipped,
                "sample": sample,
            },
            ensure_ascii=False,
            indent=2,
        )
    )


if __name__ == "__main__":
    main()
