#!/usr/bin/env python3
# -*- coding: utf-8 -*-

"""
find_qe_pseudo.py

Quantum ESPRESSO 用 UPF 擬ポテンシャル候補を検索・順位付けする。

Usage:
    python find_qe_pseudo.py PSEUDO_DIR ELEMENT [--xc auto] [--top 10]
    python find_qe_pseudo.py D:\qe\pseudo Zn
    python find_qe_pseudo.py D:\qe\pseudo O --xc r2scan
    python find_qe_pseudo.py D:\qe\pseudo Fe --xc hse06

対応する主な --xc:
    auto, r2scan, rscan, scan, pbe, pbesol, gga, meta, lda,
    hse06, hse, pbe0, mbj

方針:
  * auto    : r2SCAN > rSCAN > SCAN > PBE > PBEsol > その他GGA > LDA
  * HSE/PBE0/mBJ:
              PBE 系 pseudo を優先
  * r2SCAN  : r2SCAN > rSCAN > SCAN > PBE
  * PBE     : PBE > PBEsol > その他GGA
  * LDA     : LDA を優先

注意:
  このスクリプトの順位は「候補推薦」であり、物理的妥当性を保証するものではない。
  特に semicore の有無や相対論効果、PAW/USPP/NC の選択は用途に応じて確認すること。
"""

from __future__ import annotations

import argparse
import json
import re
import sys
from dataclasses import dataclass, asdict
from pathlib import Path
from typing import Optional


# ----------------------------------------------------------------------
# Data model
# ----------------------------------------------------------------------

@dataclass
class PseudoInfo:
    path: Path
    element: str = ""
    functional_raw: str = ""
    xc_family: str = "unknown"
    pseudo_type: str = ""
    z_valence: Optional[float] = None
    relativistic: str = ""
    ecutwfc: Optional[float] = None
    ecutrho: Optional[float] = None
    score: float = 0.0
    reasons: list[str] = None

    def __post_init__(self):
        if self.reasons is None:
            self.reasons = []


# ----------------------------------------------------------------------
# XC classification
# ----------------------------------------------------------------------

def normalize_xc_name(text: str) -> str:
    s = (text or "").lower()
    s = s.replace("_", "").replace("-", "").replace(" ", "")

    if "r2scan" in s:
        return "r2scan"
    if "rscan" in s:
        return "rscan"
    if "scan" in s:
        return "scan"
    if "pbesol" in s:
        return "pbesol"
    if "revpbe" in s:
        return "revpbe"
    if "pbe" in s:
        return "pbe"
    if "blyp" in s:
        return "blyp"
    if "pw91" in s:
        return "pw91"
    if "gga" in s:
        return "gga"
    if "lda" in s or "pz" in s or "pw" == s:
        return "lda"
    if "meta" in s or "mgga" in s:
        return "meta"
    return "unknown"


def classify_xc(functional_text: str, filename: str) -> str:
    # UPF header を優先し、不明ならファイル名を見る
    fam = normalize_xc_name(functional_text)
    if fam != "unknown":
        return fam
    return normalize_xc_name(filename)


# ----------------------------------------------------------------------
# UPF parsing
# ----------------------------------------------------------------------

def first_match(text: str, patterns: list[str]) -> str:
    for pat in patterns:
        m = re.search(pat, text, re.IGNORECASE | re.MULTILINE)
        if m:
            return m.group(1).strip()
    return ""


def first_float(text: str, patterns: list[str]) -> Optional[float]:
    s = first_match(text, patterns)
    if not s:
        return None
    try:
        return float(s.replace("D", "E").replace("d", "e"))
    except ValueError:
        return None


def parse_upf(path: Path) -> PseudoInfo:
    # UPF は通常テキスト。巨大でないので先頭 128 kB 程度を見る。
    try:
        raw = path.read_bytes()[:131072]
    except OSError:
        raw = b""

    # XML宣言がUTF-8でも、古いUPFは必ずしも厳密でない。
    text = ""
    for enc in ("utf-8", "latin-1", "cp1252"):
        try:
            text = raw.decode(enc)
            break
        except UnicodeDecodeError:
            pass
    if not text:
        text = raw.decode("latin-1", errors="replace")

    info = PseudoInfo(path=path)

    # UPF v2 attributes
    info.element = first_match(text, [
        r'\belement\s*=\s*"([^"]+)"',
        r"\belement\s*=\s*'([^']+)'",
        r'^\s*Element\s*[:=]\s*([A-Za-z]{1,3})\b',
    ])

    info.functional_raw = first_match(text, [
        r'\bfunctional\s*=\s*"([^"]+)"',
        r"\bfunctional\s*=\s*'([^']+)'",
        r'^\s*Exchange-correlation\s*[:=]\s*(.+?)\s*$',
        r'^\s*functional\s*[:=]\s*(.+?)\s*$',
    ])

    info.pseudo_type = first_match(text, [
        r'\bpseudo_type\s*=\s*"([^"]+)"',
        r"\bpseudo_type\s*=\s*'([^']+)'",
        r'^\s*Pseudo(?:potential)?\s*type\s*[:=]\s*(.+?)\s*$',
    ])

    info.relativistic = first_match(text, [
        r'\brelativistic\s*=\s*"([^"]+)"',
        r"\brelativistic\s*=\s*'([^']+)'",
    ])

    info.z_valence = first_float(text, [
        r'\bz_valence\s*=\s*"([^"]+)"',
        r"\bz_valence\s*=\s*'([^']+)'",
        r'^\s*z[_ ]?valence\s*[:=]\s*([0-9.+\-EeDd]+)',
    ])

    # 一部UPF / generator metadataに存在することがある
    info.ecutwfc = first_float(text, [
        r'\bwfc_cutoff\s*=\s*"([^"]+)"',
        r"\bwfc_cutoff\s*=\s*'([^']+)'",
        r'\bsuggested_wfc_cutoff\s*=\s*"([^"]+)"',
        r"\bsuggested_wfc_cutoff\s*=\s*'([^']+)'",
        r'^\s*(?:Suggested\s+)?cutoff.*wavefunctions.*?([0-9.+\-EeDd]+)',
    ])

    info.ecutrho = first_float(text, [
        r'\brho_cutoff\s*=\s*"([^"]+)"',
        r"\brho_cutoff\s*=\s*'([^']+)'",
        r'\bsuggested_rho_cutoff\s*=\s*"([^"]+)"',
        r"\bsuggested_rho_cutoff\s*=\s*'([^']+)'",
        r'^\s*(?:Suggested\s+)?cutoff.*charge.*?([0-9.+\-EeDd]+)',
    ])

    info.xc_family = classify_xc(info.functional_raw, path.name)
    return info


# ----------------------------------------------------------------------
# Candidate matching
# ----------------------------------------------------------------------

def normalize_element(symbol: str) -> str:
    s = symbol.strip()
    if not re.fullmatch(r"[A-Za-z]{1,3}", s):
        raise ValueError(f"invalid element symbol: {symbol!r}")
    return s[0].upper() + s[1:].lower()


def filename_matches_element(filename: str, element: str) -> bool:
    # Zn.xxx.UPF, Zn_xxx.upf, zn-xxx.UPF など
    stem = Path(filename).stem
    return bool(re.match(
        rf"^{re.escape(element)}(?:[._+\-]|$)",
        stem,
        re.IGNORECASE
    ))


def info_matches_element(info: PseudoInfo, element: str) -> bool:
    if info.element:
        if info.element.strip().lower() == element.lower():
            return True
        # Header に別元素が明記されていれば除外
        return False
    return filename_matches_element(info.path.name, element)


# ----------------------------------------------------------------------
# Ranking
# ----------------------------------------------------------------------

AUTO_XC_SCORE = {
    "r2scan": 120,
    "rscan": 110,
    "scan": 100,
    "pbe": 90,
    "pbesol": 82,
    "revpbe": 76,
    "pw91": 74,
    "blyp": 72,
    "gga": 70,
    "meta": 65,
    "lda": 30,
    "unknown": 0,
}

REQUESTED_XC_SCORE = {
    "r2scan": {
        "r2scan": 140, "rscan": 120, "scan": 110,
        "pbe": 70, "pbesol": 55, "gga": 45, "lda": 0,
    },
    "rscan": {
        "rscan": 140, "r2scan": 125, "scan": 115,
        "pbe": 70, "gga": 45, "lda": 0,
    },
    "scan": {
        "scan": 140, "r2scan": 125, "rscan": 120,
        "pbe": 70, "gga": 45, "lda": 0,
    },
    "pbe": {
        "pbe": 140, "pbesol": 95, "revpbe": 85,
        "gga": 80, "pw91": 75, "blyp": 65,
        "r2scan": 45, "scan": 40, "lda": 0,
    },
    "pbesol": {
        "pbesol": 140, "pbe": 105, "gga": 85,
        "r2scan": 50, "scan": 45, "lda": 0,
    },
    "gga": {
        "pbe": 130, "pbesol": 120, "revpbe": 110,
        "pw91": 105, "blyp": 100, "gga": 95,
        "r2scan": 50, "scan": 45, "lda": 0,
    },
    "meta": {
        "r2scan": 140, "rscan": 130, "scan": 125,
        "meta": 115, "pbe": 65, "gga": 50, "lda": 0,
    },
    "lda": {
        "lda": 140, "pbe": 20, "gga": 15, "unknown": 0,
    },

    # Hybrid / mBJ は通常 PBE 系 pseudo を優先
    "hse06": {
        "pbe": 140, "pbesol": 90, "gga": 80,
        "r2scan": 45, "scan": 40, "lda": 0,
    },
    "hse": {
        "pbe": 140, "pbesol": 90, "gga": 80,
        "r2scan": 45, "scan": 40, "lda": 0,
    },
    "pbe0": {
        "pbe": 140, "pbesol": 90, "gga": 80,
        "r2scan": 45, "scan": 40, "lda": 0,
    },
    "mbj": {
        "pbe": 140, "pbesol": 95, "gga": 80,
        "r2scan": 45, "scan": 40, "lda": 0,
    },
}


def pseudo_type_bonus(info: PseudoInfo) -> tuple[float, list[str]]:
    """
    形式だけで「PAW > USPP > NC」と物理的優劣を決めない。
    ただし判別可能・一般的な系列であることに小さな加点を与える。
    """
    name = (info.path.name + " " + info.pseudo_type).lower()
    score = 0.0
    reasons = []

    if "paw" in name or "kjpaw" in name:
        score += 8
        reasons.append("PAW")
    elif "uspp" in name or "ultrasoft" in name or re.search(r"\bus\b", name):
        score += 6
        reasons.append("USPP")
    elif "oncv" in name or "norm-conserv" in name or re.search(r"\bnc\b", name):
        score += 6
        reasons.append("NC/ONCV")

    # よく使われる配布系列。ただし品質保証ではない。
    if "sssp" in name:
        score += 10
        reasons.append("SSSP")
    elif "psl" in name or "pslibrary" in name:
        score += 6
        reasons.append("PSLibrary")
    elif "pseudo-dojo" in name or "dojo" in name:
        score += 6
        reasons.append("PseudoDojo")

    return score, reasons


def feature_bonus(info: PseudoInfo) -> tuple[float, list[str]]:
    name = info.path.name.lower()
    score = 0.0
    reasons = []

    # semicore を表すことが多いファイル名。用途依存なので小さな加点に留める。
    semicore_tokens = ("spn", "dn", "sp", "semicore", "_sv", "-sv")
    if any(tok in name for tok in semicore_tokens):
        score += 5
        reasons.append("semicore-like")

    if info.z_valence is not None:
        score += 2
        reasons.append(f"Zval={info.z_valence:g}")

    if info.ecutwfc is not None or info.ecutrho is not None:
        score += 2
        reasons.append("cutoff metadata")

    if info.relativistic:
        rel = info.relativistic.lower()
        if rel not in ("no", "none", "false"):
            score += 1
            reasons.append(f"rel={info.relativistic}")

    return score, reasons


def score_pseudo(info: PseudoInfo, requested_xc: str) -> PseudoInfo:
    xc = requested_xc.lower()

    if xc == "auto":
        base = AUTO_XC_SCORE.get(info.xc_family, 0)
    else:
        table = REQUESTED_XC_SCORE.get(xc, {})
        base = table.get(info.xc_family, 0)

    info.score = float(base)
    if base:
        info.reasons.append(f"XC={info.xc_family}")

    b, r = pseudo_type_bonus(info)
    info.score += b
    info.reasons.extend(r)

    b, r = feature_bonus(info)
    info.score += b
    info.reasons.extend(r)

    # header で元素を明示できているものを少し優先
    if info.element:
        info.score += 2
        info.reasons.append("element metadata")

    return info


# ----------------------------------------------------------------------
# Search / output
# ----------------------------------------------------------------------

def find_candidates(pseudo_dir: Path, element: str, requested_xc: str) -> list[PseudoInfo]:
    files = sorted([
        p for p in pseudo_dir.rglob("*")
        if p.is_file() and p.suffix.lower() == ".upf"
    ])

    result = []
    for p in files:
        # まず filename で粗く絞る。Header の element があれば後で確定。
        if not filename_matches_element(p.name, element):
            # ファイル名が独特なケースもあるので header を読む
            info = parse_upf(p)
            if not info_matches_element(info, element):
                continue
        else:
            info = parse_upf(p)
            if not info_matches_element(info, element):
                continue

        result.append(score_pseudo(info, requested_xc))

    result.sort(
        key=lambda x: (
            -x.score,
            x.xc_family,
            x.path.name.lower()
        )
    )
    return result


def fmt_float(v: Optional[float]) -> str:
    return "-" if v is None else f"{v:g}"


def print_table(items: list[PseudoInfo], top: int, requested_xc: str):
    shown = items[:top] if top > 0 else items

    print(f"Requested XC : {requested_xc}")
    print(f"Candidates   : {len(items)}")
    print()

    if not shown:
        print("No matching UPF file found.")
        return

    for i, info in enumerate(shown, 1):
        mark = "  <-- recommended" if i == 1 else ""
        print(f"[{i:2d}] score={info.score:6.1f}  {info.path.name}{mark}")
        print(f"     XC       : {info.xc_family}"
              + (f"  ({info.functional_raw})" if info.functional_raw else ""))
        print(f"     type     : {info.pseudo_type or '-'}")
        print(f"     Zval     : {fmt_float(info.z_valence)}")
        print(f"     ecutwfc  : {fmt_float(info.ecutwfc)}")
        print(f"     ecutrho  : {fmt_float(info.ecutrho)}")
        print(f"     reasons  : {', '.join(info.reasons) if info.reasons else '-'}")
        print(f"     path     : {info.path}")
        print()


def main():
    parser = argparse.ArgumentParser(
        description="Search and rank Quantum ESPRESSO UPF pseudopotentials."
    )
    parser.add_argument("pseudo_dir", help="Pseudo directory")
    parser.add_argument("element", help="Element symbol, e.g. Zn, O, Fe")
    parser.add_argument(
        "--xc",
        default="auto",
        choices=[
            "auto", "r2scan", "rscan", "scan",
            "pbe", "pbesol", "gga", "meta", "lda",
            "hse06", "hse", "pbe0", "mbj"
        ],
        help="Target XC functional (default: auto)"
    )
    parser.add_argument(
        "--top", type=int, default=10,
        help="Number of candidates to display; 0 = all (default: 10)"
    )
    parser.add_argument(
        "--json", action="store_true",
        help="Output JSON instead of text"
    )
    args = parser.parse_args()

    try:
        element = normalize_element(args.element)
    except ValueError as e:
        parser.error(str(e))

    pseudo_dir = Path(args.pseudo_dir).expanduser()
    if not pseudo_dir.is_dir():
        parser.error(f"pseudo directory not found: {pseudo_dir}")

    items = find_candidates(pseudo_dir, element, args.xc)

    if args.json:
        data = []
        for x in (items[:args.top] if args.top > 0 else items):
            d = asdict(x)
            d["path"] = str(x.path)
            data.append(d)
        print(json.dumps(data, ensure_ascii=False, indent=2))
    else:
        print(f"Element      : {element}")
        print(f"Pseudo dir   : {pseudo_dir.resolve()}")
        print_table(items, args.top, args.xc)


if __name__ == "__main__":
    main()
