#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
cif_sqlite_index.py

再帰的に CIF を検索して SQLite3 DB を作成し、組成から検索するツール。

例:
  python search_cif_db.py --mode index  --root ./COD --db cif_index.sqlite
  python search_cif_db.py --mode search --db cif_index.sqlite --formula BaTiO3
  python search_cif_db.py --mode search --db cif_index.sqlite --formula BaTiO3 --match exact
  python search_cif_db.py --mode search --db cif_index.sqlite --formula BaTiO3 --match contains
"""

from __future__ import annotations

import argparse
import json
import re
import warnings
import sqlite3
import sys
import traceback
from pathlib import Path
from typing import Any


try:
    from pymatgen.core import Composition, Structure
except Exception:
    print("ERROR: failed to import pymatgen")
    traceback.print_exc()
    print("\nInstall example:")
    print("  pip install pymatgen")
    input("\nPress ENTER to terminate>>\n")
    sys.exit(1)


warnings.simplefilter("ignore")


def guess_db_name(path: Path) -> str:
    """パスから cod/tcod/pcod などを大雑把に推定する。"""
    parts = [p.lower() for p in path.parts]
    for name in ("pcod", "tcod", "cod"):
        if name in parts:
            return name
    return "unknown"


def guess_material_id(path: Path) -> str:
    """ファイル名から数字IDを推定。例: 1234567.cif -> 1234567"""
    m = re.search(r"(\d+)", path.stem)
    return m.group(1) if m else path.stem


def normalize_formula(formula: str) -> str:
    """BaTiO3, Ba Ti O3 などを pymatgen Composition で標準化。"""
    comp = Composition(formula)
    return comp.reduced_formula


def composition_info_from_structure(path: Path) -> dict[str, Any]:
    """CIFからStructureを読み、検索用メタデータを作る。"""
    s = Structure.from_file(str(path))

    comp = s.composition
    red_comp = comp.reduced_composition

    elements = sorted([el.symbol for el in red_comp.elements])
    element_set = ",".join(elements)

    try:
        sg_symbol = s.get_space_group_info()[0]
        sg_number = int(s.get_space_group_info()[1])
    except Exception:
        sg_symbol = ""
        sg_number = -1

    return {
        "formula": comp.formula,
        "reduced_formula": red_comp.reduced_formula,
        "anonymous_formula": red_comp.anonymized_formula,
        "elements": element_set,
        "nelements": len(elements),
        "nsites": len(s),
        "sg_symbol": sg_symbol,
        "sg_number": sg_number,
        "volume": float(s.volume),
        "a": float(s.lattice.a),
        "b": float(s.lattice.b),
        "c": float(s.lattice.c),
        "alpha": float(s.lattice.alpha),
        "beta": float(s.lattice.beta),
        "gamma": float(s.lattice.gamma),
    }


def create_schema(conn: sqlite3.Connection) -> None:
    conn.execute("""
    CREATE TABLE IF NOT EXISTS cif_index (
        id INTEGER PRIMARY KEY AUTOINCREMENT,
        source_db TEXT,
        material_id TEXT,
        cif_path TEXT UNIQUE,
        formula TEXT,
        reduced_formula TEXT,
        anonymous_formula TEXT,
        elements TEXT,
        nelements INTEGER,
        nsites INTEGER,
        sg_symbol TEXT,
        sg_number INTEGER,
        volume REAL,
        a REAL,
        b REAL,
        c REAL,
        alpha REAL,
        beta REAL,
        gamma REAL,
        status TEXT,
        error TEXT
    )
    """)

    conn.execute("CREATE INDEX IF NOT EXISTS idx_formula ON cif_index(reduced_formula)")
    conn.execute("CREATE INDEX IF NOT EXISTS idx_elements ON cif_index(elements)")
    conn.execute("CREATE INDEX IF NOT EXISTS idx_source_db ON cif_index(source_db)")
    conn.commit()


def upsert_record(conn: sqlite3.Connection, rec: dict[str, Any]) -> None:
    conn.execute("""
    INSERT OR REPLACE INTO cif_index (
        source_db, material_id, cif_path,
        formula, reduced_formula, anonymous_formula,
        elements, nelements, nsites,
        sg_symbol, sg_number,
        volume, a, b, c, alpha, beta, gamma,
        status, error
    ) VALUES (
        :source_db, :material_id, :cif_path,
        :formula, :reduced_formula, :anonymous_formula,
        :elements, :nelements, :nsites,
        :sg_symbol, :sg_number,
        :volume, :a, :b, :c, :alpha, :beta, :gamma,
        :status, :error
    )
    """, rec)


def build_index(
    root: Path,
    db_path: Path,
    pattern: str = "*.cif",
    store_errors: int = 1,
    commit_interval: int = 200,
) -> None:
    conn = sqlite3.connect(str(db_path))
    create_schema(conn)

    files = list(root.rglob(pattern))
    print(f"Found CIF files: {len(files)}")

    n_ok = 0
    n_err = 0

    for i, path in enumerate(files, start=1):
        path = path.resolve()

        rec_base = {
            "source_db": guess_db_name(path),
            "material_id": guess_material_id(path),
            "cif_path": str(path),
        }

        try:
            info = composition_info_from_structure(path)
            rec = {
                **rec_base,
                **info,
                "status": "ok",
                "error": "",
            }
            upsert_record(conn, rec)
            n_ok += 1

        except Exception as e:
            n_err += 1
            print(f"[ERROR] {path}: {e}")

            if store_errors:
                rec = {
                    **rec_base,
                    "formula": "",
                    "reduced_formula": "",
                    "anonymous_formula": "",
                    "elements": "",
                    "nelements": -1,
                    "nsites": -1,
                    "sg_symbol": "",
                    "sg_number": -1,
                    "volume": -1.0,
                    "a": -1.0,
                    "b": -1.0,
                    "c": -1.0,
                    "alpha": -1.0,
                    "beta": -1.0,
                    "gamma": -1.0,
                    "status": "error",
                    "error": str(e),
                }
                upsert_record(conn, rec)

        if i % commit_interval == 0:
            conn.commit()
            print(f"Indexed {i}/{len(files)}  ok={n_ok}  error={n_err}")

    conn.commit()
    conn.close()

    print(f"Done. ok={n_ok}, error={n_err}")
    print(f"DB: {db_path}")


def search_by_formula(
    db_path: Path,
    formula: str,
    source_db: str = "",
    match: str = "exact",
    limit: int = 100,
    output_json: int = 0,
) -> None:
    target_comp = Composition(formula).reduced_composition
    target_formula = target_comp.reduced_formula
    target_elements = sorted([el.symbol for el in target_comp.elements])
    target_element_set = ",".join(target_elements)

    conn = sqlite3.connect(str(db_path))
    conn.row_factory = sqlite3.Row

    params: list[Any] = []
    where = ["status = 'ok'"]

    if source_db:
        where.append("source_db = ?")
        params.append(source_db)

    if match == "exact":
        # BaTiO3 と TiBaO3 はどちらも BaTiO3 に正規化される
        where.append("reduced_formula = ?")
        params.append(target_formula)

    elif match == "elements":
        # 元素集合が完全一致：Ba-Ti-O系のみ
        # BaTiO3, Ba2TiO4, BaTi2O5 はヒット
        # TiO2, BaO, Ba, Ti はヒットしない
        where.append("elements = ?")
        params.append(target_element_set)

    elif match == "contains":
        # 指定元素の任意の部分集合を許す
        # BaTiO3検索で Ba, Ti, TiO2, BaO などもヒット
        #
        # DB側 elements は "Ba,O,Ti" のようなカンマ区切りなので、
        # "," || elements || "," にして完全な元素名単位で検索する
        subset_conditions = []
        for el in target_elements:
            subset_conditions.append("? LIKE '%,' || elements || ',%'")
        where.append("(" + " OR ".join(subset_conditions) + ")")
        params.extend(["," + target_element_set + ","] * len(target_elements))

    else:
        raise ValueError(f"Unknown match mode: {match}")

    sql = f"""
    SELECT
        source_db, material_id, reduced_formula, formula,
        elements, sg_symbol, sg_number, nsites,
        a, b, c, alpha, beta, gamma, volume,
        cif_path
    FROM cif_index
    WHERE {' AND '.join(where)}
    ORDER BY nelements DESC, source_db, reduced_formula, material_id
    LIMIT ?
    """
    params.append(limit)

    rows = [dict(r) for r in conn.execute(sql, params).fetchall()]
    conn.close()

    if output_json:
        print(json.dumps(rows, ensure_ascii=False, indent=2))
        return

    print(f"query formula   : {formula}")
    print(f"reduced formula : {target_formula}")
    print(f"elements        : {target_element_set}")
    print(f"match           : {match}")
    print(f"hits            : {len(rows)}")
    print()

    for r in rows:
        print(
            f"[{r['source_db']}] {r['material_id']}  "
            f"{r['reduced_formula']}  "
            f"elements={r['elements']}  "
            f"SG={r['sg_symbol']}({r['sg_number']})  "
            f"nsites={r['nsites']}"
        )
        print(f"  a,b,c = {r['a']:.6g}, {r['b']:.6g}, {r['c']:.6g}")
        print(f"  path  = {r['cif_path']}")
        print()


def show_info(db_path: Path) -> None:
    conn = sqlite3.connect(str(db_path))
    cur = conn.cursor()

    print("Total:")
    for row in cur.execute("""
        SELECT status, COUNT(*) FROM cif_index GROUP BY status
    """):
        print(f"  {row[0]}: {row[1]}")

    print("\nBy source_db:")
    for row in cur.execute("""
        SELECT source_db, status, COUNT(*)
        FROM cif_index
        GROUP BY source_db, status
        ORDER BY source_db, status
    """):
        print(f"  {row[0]:8s} {row[1]:8s} {row[2]}")

    conn.close()


def main() -> None:
    parser = argparse.ArgumentParser()

    parser.add_argument("--mode", type=str, default="search",
                        choices=["index", "search", "info"])

    parser.add_argument("--root", type=str, default=".")
    parser.add_argument("--db", type=str, default="cif_index.sqlite")
    parser.add_argument("--pattern", type=str, default="*.cif")

    parser.add_argument("--formula", type=str, default="")
    parser.add_argument("--match", type=str, default="exact",
                    choices=["exact", "elements", "contains"],
                    help=(
                        "exact: 還元組成一致, "
                        "elements: 元素集合一致, "
                        "contains: 指定元素集合の部分集合も許す"
                    ))
    parser.add_argument("--source-db", type=str, default="",
                        help="cod, tcod, pcod など。空なら全DB検索。")

    parser.add_argument("--allow-subset", type=int, default=0, choices=[0, 1],
                        help="1なら組成式完全一致ではなく、元素集合一致で検索する。")
    parser.add_argument("--store-errors", type=int, default=1, choices=[0, 1])
    parser.add_argument("--json", type=int, default=0, choices=[0, 1])
    parser.add_argument("--limit", type=int, default=100)

    args = parser.parse_args()

    db_path = Path(args.db)

    if args.mode == "index":
        build_index(
            root=Path(args.root),
            db_path=db_path,
            pattern=args.pattern,
            store_errors=args.store_errors,
        )

    elif args.mode == "search":
        if not args.formula:
            print("ERROR: --formula is required for --mode search")
            sys.exit(1)

        search_by_formula(
            db_path=db_path,
            formula=args.formula,
            source_db=args.source_db,
            match=args.match,
            limit=args.limit,
            output_json=args.json,
        )

    elif args.mode == "info":
        show_info(db_path)


if __name__ == "__main__":
    main()