from win32com.client import Dispatch
import matplotlib.pyplot as plt
import numpy as np
from scipy.optimize import fmin
import pandas as pd

#Granodiorite in Type III, Ra surface complexation and precipitation

def isotherm(input_param, test=True):
    iphreeqc = Dispatch("IPhreeqcCOM.Object")
    iphreeqc.LoadDatabase("PHREEQC_Davies_e-_ThermoChimie_v12a.dat")
    
    ###################################################################
    # INPUT parameters   
    X1_coeff=input_param[0]
    X2_coeff=input_param[1]
    Ksp=input_param[2]
    

    S_conc=np.array([2.030e-3, 2.030e-3, 2.030e-3, 2.030e-3, 2.030e-3])
    Ra_conc=np.array([4.689E-09,4.296E-09,4.158E-09,4.298E-09,3.751E-09])
    Ba_conc=np.array([1.132E-03,1.123E-04,1.211E-05,1.772E-06,6.569E-07])
    # Ksp=-10.25

    result1=[[0 for col in range(1)] for row in range(5)]   # only row numbers should be defined correctly
    result2=[[0 for col in range(1)] for row in range(5)]   # only row numbers should be defined correctly
        
    for i in range(0,len(S_conc)):
        str_ini_part1="""
       SURFACE_MASTER_SPECIES
        S_w           S_wOH        
         
        SURFACE_SPECIES
        S_wOH = S_wOH
        log_k     0
        H+ + S_wOH = S_wOH2+
        log_k     5.0491446
        S_wOH = S_wO- + H+
        log_k     -8.78103445

        S_wOH + Ba+2 = S_wOBa+ + H+        
        log_k     """
        str_ini_part1_1=str(X1_coeff)
        str_ini_part1_2="""   
        S_wOH + Ra+2 = S_wORa+ + H+
        log_k    """
        str_ini_part1_3=str(X2_coeff)
        str_ini_part1_4="""
        PHASES
        Barite
            Ba(SO4) = +1.000Ba+2     +1.000SO4-2  
            log_k   -9.97
        Ra(SO4)(s)
            Ra(SO4) = +1.000Ra+2     +1.000SO4-2    
            log_k   """
        str_ini_part2=str(Ksp)
        str_ini_part3="""
        Fix_H+
            H+ = H+
            log_k     0 
        END
        """
        str_ini_part4="""
        SOLUTION  1
        temp      25
        pH        7.5
        pe        4
        redox     pe
        units     mol/L
        density   1
        K         3.581e-4
        Cs        5.267e-9
        Mg        2.181e-3
        Ca        3.668e-2
        Sr        1.940e-4
        C(4)      3.606e-4
        S(-2)     9.356e-7
        Cl        1.451e-1    charge
        Na        7.351e-2
        S(6)      """
        str_ini_part5=str(S_conc[i])
        str_ini_part6="""   
        Ra        """
        str_ini_part7=str(Ra_conc[i])
        str_ini_part8="""
        Ba        """
        str_ini_part9=str(Ba_conc[i])    
        str_ini_part10="""
        -water    0.01 # kg
        
        SOLID_SOLUTIONS 1
        BaRaSO4
        -comp  Ra(SO4)(s)      0    
        -comp  Barite          0
        SAVE SOLUTION 1
        
        SELECTED_OUTPUT
        -reset                false
        USER_PUNCH
        -Heading 
        10 PUNCH SUM_S_S("BaRaSO4", "Ba")
        20 PUNCH SUM_S_S("BaRaSO4", "Ra")
        30 PUNCH TOTMOLE("Ba")  # Totol moles only in solution
        40 PUNCH TOTMOLE("Ra")  # Totol moles only in solution
        60 PUNCH SYS("Ba")  # moles in all system
        70 PUNCH SYS("Ra")  # moles in all system
        END
        
        USE SOLUTION 1
        SURFACE 1
        -sites DENSITY
        S_wOH      5.6       0.106      0.5
        
        SELECTED_OUTPUT 
        -reset                false
        USER_PUNCH
        -Heading 
        10 PUNCH SUM_S_S("BaRaSO4", "Ba")
        20 PUNCH SUM_S_S("BaRaSO4", "Ra")
        30 PUNCH TOTMOLE("Ba")  # Totol moles only in solution
        40 PUNCH TOTMOLE("Ra")  # Totol moles only in solution
        60 PUNCH SYS("Ba")  # moles in all system
        70 PUNCH SYS("Ra")  # moles in all system
        END 
        """
        input_str=str_ini_part1+str_ini_part1_1+str_ini_part1_2+str_ini_part1_3+str_ini_part1_4+str_ini_part2+str_ini_part3+str_ini_part4+str_ini_part5+str_ini_part6+str_ini_part7+str_ini_part8+str_ini_part9+str_ini_part10
        #print(input_str)
        iphreeqc.RunString(input_str)
        run_result=iphreeqc.GetSelectedOutputArray()
        print(run_result)
        
        if i==0:
            result1[i]=[entry for entry in run_result][2:3]
        else:
            result1[i]=[entry for entry in run_result][2:3]
        
        if i==0:
            result2[i]=[entry for entry in run_result][3:]
        else:
            result2[i]=[entry for entry in run_result][3:]
        # print(result)
    
    result_all1=sum(result1, [])
    result_all1=np.array(result_all1)
    # print(result_all1)

    result_all2=sum(result2, [])
    result_all2=np.array(result_all2)
    # print(result_all2)
    
    result_Ba_SS=result_all1[:,0]
    result_Ra_SS=result_all1[:,1]
    result_Ba_aqueous=result_all2[:,2]
    result_Ra_aqueous=result_all2[:,3]
    result_Ba_all=result_all1[:,4]
    result_Ra_all=result_all1[:,5]

    #qe_model=result_Ra_SS*226*1000/0.5
    qe_model=(result_Ra_all-result_Ra_aqueous)*226*1000/0.5
    qe_exp=np.array([2.14E-05,1.90E-05,1.78E-05,1.75E-05,1.40E-05])

    if test:
        print(result_all1)
        print(result_all2)
        plt.figure()
        plt.plot(Ba_conc, qe_model, "o")
        plt.plot(Ba_conc, qe_exp, "s")
        plt.xscale("log")
        print(qe_model)

    square_error=((qe_model-qe_exp)*1000)**2
    sum_square_error=np.sum(square_error)
    return sum_square_error
    
start_para=np.array([5,-1.5, -10.25])
para_optimized1=fmin(isotherm, start_para)
print(para_optimized1)

#para_optimized1=np.array([6.66044254, -1.08479668, -9.15468886])
isotherm(para_optimized1, test=True)