import pandas as pd
import matplotlib.pyplot as plt
import sys
import re
import argparse
import numpy as np

def clean_column(name):
    return re.sub(r'\[.*?\]', '', name).strip()

def plot_sorted(ax, x, y, xlabel, ylabel, marker, xscale='linear', yscale='linear'):
    x_sorted, y_sorted = zip(*sorted(zip(x, y)))
    ax.plot(x_sorted, y_sorted, marker=marker)
    ax.set_xlabel(xlabel)
    ax.set_ylabel(ylabel)
    ax.set_xscale(xscale)
    ax.set_yscale(yscale)
    ax.grid(True)

def main(infile, T_target, eps):
    # 1行目のラベルを取得して DataFrame に読み込む
    with open(infile, 'r') as f:
        header = f.readline().strip().lstrip('#').split()
    df = pd.read_csv(infile, delim_whitespace=True, comment='#', names=header)

    # 列名をクリーンアップ
    original_columns = df.columns.tolist()
    cleaned_columns = [clean_column(col) for col in original_columns]
    df.columns = cleaned_columns

    print("\nOriginal labels:", original_columns)
    print("Cleaned labels :", df.columns.tolist())

    # temperature列があるか確認
    if 'temperature' not in df.columns:
        print("Error: 'temperature' column not found in the data.")
        sys.exit(1)

    # 指定温度 ± eps の範囲で抽出
    mask = np.abs(df['temperature'] - T_target) < eps
    df_T = df[mask]

    if df_T.empty:
        print(f"\nNo data found within ±{eps} K of T = {T_target} K.")
        sys.exit(1)

    print(f"\nSelected temperatures: {df_T['temperature'].unique()}")

    # プロット設定（基本4枚）
    fig, axs = plt.subplots(2, 2, figsize=(10, 8))
    fig.suptitle(f"Transport Properties near T = {T_target} K (±{eps} K)", fontsize=14)

    # Fermi_level vs doping
    plot_sorted(axs[0, 0], df_T['doping'], abs(df_T['doping']), marker='o',
                xlabel='Fermi level [eV]', ylabel='Doping [cm^3]', xscale='linear', yscale='log')
    axs[0, 0].set_title('Doping vs Fermi Level')

    # doping vs cond_xx
    plot_sorted(axs[0, 1], df_T['Fermi_level'], abs(df_T['cond_xx']), marker='o',
                xlabel='Fermi Level [eV]', ylabel='Conductivity xx [S/m]', xscale='linear', yscale='log')
    axs[0, 1].set_title('Doping vs cond_xx')

    # Fermi_level vs seebeck_xx
    plot_sorted(axs[1, 0], df_T['Fermi_level'], abs(df_T['seebeck_xx']), marker='o',
                xlabel='Fermi Level [eV]', ylabel='Seebeck xx [µV/K]', xscale='linear', yscale='log')
    axs[1, 0].set_title('Fermi Level vs Seebeck xx')

    # Fermi_level vs seebeck_zz
    plot_sorted(axs[1, 1], df_T['Fermi_level'], abs(df_T['seebeck_zz']), marker='o',
                xlabel='Fermi Level [eV]', ylabel='Seebeck zz [µV/K]', xscale='linear', yscale='log')
    axs[1, 1].set_title('Fermi Level vs Seebeck zz')

    plt.tight_layout(rect=[0, 0.03, 1, 0.95])
    plt.show()

    # --- Mottの式による D(E) 推定と追加プロット ---
    # 定数（eV単位）
    kB_eV = 8.617333262e-5  # eV/K
    coeff_eV = -3 / (np.pi**2 * kB_eV**2 * T_target)  # [1/eV]

    # SeebeckからdlnD/dEを計算（µV/K → V/K）
    df_T['dlnDxx_dE'] = coeff_eV * df_T['seebeck_xx'] * 1e-6
    df_T['dlnDzz_dE'] = coeff_eV * df_T['seebeck_zz'] * 1e-6

    # 積分でD(E)を推定（初期値D0=1として相対値）
    df_T_sorted = df_T.sort_values(by='Fermi_level')
    EF_vals = df_T_sorted['Fermi_level'].values
    dlnDxx_vals = df_T_sorted['dlnDxx_dE'].values
    dlnDzz_vals = df_T_sorted['dlnDzz_dE'].values

    Dxx = [1.0]
    Dzz = [1.0]
    for i in range(1, len(EF_vals)):
        dE = EF_vals[i] - EF_vals[i - 1]  # [eV]
        Dxx.append(Dxx[-1] * np.exp(dlnDxx_vals[i - 1] * dE))
        Dzz.append(Dzz[-1] * np.exp(dlnDzz_vals[i - 1] * dE))

    # グラフ追加（新しいFigure）
    fig2, axs2 = plt.subplots(1, 2, figsize=(12, 4))
    fig2.suptitle(f"Estimated D(E) from Mott relation at T = {T_target} K", fontsize=14)

    axs2[0].plot(EF_vals, Dxx, marker='o')
    axs2[0].set_xlabel('Fermi Level [eV]')
    axs2[0].set_ylabel('Relative Dxx(E)')
    axs2[0].set_title('Fermi Level vs Dxx(E)')
    axs2[0].grid(True)

    axs2[1].plot(EF_vals, Dzz, marker='o', color='tab:red')
    axs2[1].set_xlabel('Fermi Level [eV]')
    axs2[1].set_ylabel('Relative Dzz(E)')
    axs2[1].set_title('Fermi Level vs Dzz(E)')
    axs2[1].grid(True)

    plt.tight_layout(rect=[0, 0.03, 1, 0.95])
    plt.show()

if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("-i", "--infile", help="Input transport data file")
    parser.add_argument("--T", type=float, default=300.0, help="Target temperature in K")
    parser.add_argument("--eps", type=float, default=1.0, help="Tolerance for temperature matching")
    args = parser.parse_args()
    main(args.infile, args.T, args.eps)
