#exec(open("BolsasAnalises.py").read())

import gudhi
import gudhi.representations
import numpy as np
import matplotlib.pyplot as plt
import datetime as datetime
import time
import matplotlib.dates as mdates
from scipy.spatial.distance import pdist, squareform
import random
import os
import glob
import sys
import warnings

exec(open("Bolsas.py").read())
exec(open("BolsasAux.py").read())
#exec(open("auxMatlabEx.py").read())

if sys.platform == "win32":
    exec(open("smeUtils.py").read())
    exec(open("smeGudhi.py").read())
else:
    from smeUtils import *
    from smeGudhi import *


#dname = "/db/Economia/DataOrg2026_dias010/"
#      = "/db/Economia/DataOrg2026_318_qt10_st1/NY_19971001_19971014__318_dim9_qt10_st1.txt"

# ba1 = bolsasAnalisa("/db/Economia/", '19900101', '19900131', edgeMax=0.2, persLim=0.03, cols=[1,2,3], maxDim=4, lenInt=10, stInt=1, nEmp=318)
# ba = bolsasAnalisa("/db/Economia/", '19890421', '20090420', edgeMax=0.2, persLim=0.03, cols=[1,2,3], maxDim=4, lenInt=10, stInt=1, nEmp=318)
class bolsasAnalisa:
    def __init__ (self, dname, date1, date2, lenInt=10, stInt=1, edgeMax=0.2, persLim=0.03, cols=[1,2,3], maxDim=4, nEmp=318):
        self.date1 = date1
        self.date2 = date2
        self.edgeMax = edgeMax
        self.persLim = persLim
        
        self.lenInt = lenInt
        self.stInt= stInt
        self.nEmp = nEmp
        self.vData = []
        self.qtError = 0

        realTime = time.time()
        cpuTime = time.process_time()

        #if dname[-1] != os.sep: dname = dname + os.sep
        #self.dname = f"{dname}DataOrg2026_{nEmp}_qt{lenInt}_st{stInt}{os.sep}"
        if dname[-1] != "/": dname = dname + "/"
        self.dname = f"{dname}DataOrg2026_{nEmp}_qt{lenInt}_st{stInt}/"
        if not os.path.isdir(self.dname):
            print(f"\n***** Erro: a pasta {self.dname} não existe!\n")
            return
        
        self.vFiles = sorted(glob.glob(self.dname + "*.txt"))

        qtFiles = len(self.vFiles)
        if qtFiles == 0:
            print(f"\n     ****** Não encontrou ficheiros")
            print(f"     Pasta de pesquisa: {self.dname}\n")
            return

        n1 = -1
        for fname in self.vFiles:
            n1 += 1
            fname1 = os.path.splitext(os.path.basename(fname))[0]
            dia1 = fname1[3:11]
            dia2 = fname1[12:20]
            if dia1 > date2: break
            if dia1 < date1: continue

            rc = BolsasFile(fname, edgeMax=edgeMax, maxDim=maxDim, sparse=True, skip=0, cols=cols, persLim=persLim, verbose=0)
            print(f"\rDia inicial: {dia1[:4]}-{dia1[4:6]}-{dia1[6:8]}", end="", flush=True)
            if rc.error < 0:
                self.qtError += 1
                continue
            nBetti = [0,0,0,0,0]
            n2 = 0
            len2 = len(rc.numBetti)
            while n2 < len2 and n2 < 5:
                nBetti[n2] = rc.numBetti[n2]
                n2 += 1

            info1 = {'fname':fname, 'diaI':dia1, 'diaF':dia2, 'lenInt':self.lenInt, stInt:self.stInt, 
                     'vol':rc.vol, 'nBetti':nBetti, 'nHoles':len(rc.holes), 'holes':rc.holes} # , 'pers':rc.persistence}
            self.vData.append(info1)
            n1 += 1
        print(f"\rLeu {n1} ficheiros da pasta: '{self.dname}'")
        print(f"realTempo: {(time.time() - realTime):.2f}  ;  cpuTime: {(time.process_time() - cpuTime):.2f}")


# Primeiro correr o "bolsasAnalisa" para fazer o varrimento e guardá-lo no objecto "ba"
#     ba1 = bolsasAnalisa("/db/Economia/", '19890101', '20101231', edgeMax=0.2, persLim=0.03, cols=[1,2,3], maxDim=4)
# Depois fazer o gráfico que se deseja (exemplo):
#     graphAnalisa (ba1, '19890101', '20101231', graphs=['V','B0','H'])
# Valores válidos para "graphs" (para já): "V", "B0", "B1", "B2", "B3", "B4", "H"
def graphAnalisa (ba, diaI, diaF, graphs=['vol','B0'], volNorm=True, centra=False):
    import matplotlib.pyplot as plt
    import matplotlib.dates as mdates

    x = []
    xd = []
    vol = []
    nBetti0 = []
    nBetti1 = []
    nBetti2 = []
    nBetti3 = []
    nBetti4 = []
    nHoles = []
    n1 = 0

    diaI1 = diaI[0:4] + "-" + diaI[4:6] + "-" + diaI[6:8]
    diaF1 = diaF[0:4] + "-" + diaF[4:6] + "-" + diaF[6:8]
    t1 = time.mktime(time.strptime(diaI1, "%Y-%m-%d"))
    t2 = time.mktime(time.strptime(diaF1, "%Y-%m-%d"))
    qtDias = int((t2-t1)/86400)
    
    print(f"DiaI: {diaI} ; DiaF: {diaF} ; len: {len(ba.vData)}")
    for p1 in ba.vData:
        #print(f"__p1_DiaI: {p1['diaI']} ; __p1_DiaF: {p1['diaF']}")
        if int(p1['diaI']) > int(diaF): break
        if int(p1['diaI']) < int(diaI): continue
        x.append(n1)
        xd.append(p1['diaI'][0:4] + '/' + p1['diaI'][4:6] + '/' + p1['diaI'][6:8])
        vol.append(p1['vol'])
        nBetti0.append(p1['nBetti'][0])
        nBetti1.append(p1['nBetti'][1])
        nBetti2.append(p1['nBetti'][2])
        nBetti3.append(p1['nBetti'][3])
        nBetti4.append(p1['nBetti'][4])
        nHoles.append(p1['nHoles'])
        n1 += 1
    print(f"Número de pontos do gráfico: {len(x)}")

    if n1 < 1:
        print("***** Não existem pontos...")
        return
    
    xDatas = [datetime.datetime.strptime(d,'%Y/%m/%d').date() for d in xd]
    #y = range(len(xd))

    #print(f"Len(xd): {len(xd)}")

    maxBetti0 = max(nBetti0)
    maxBetti1 = max(nBetti1)
    maxBetti2 = max(nBetti2)
    maxBetti3 = max(nBetti3)
    maxBetti4 = max(nBetti4)
    maxHoles = max(nHoles)
    maxVol = max(vol)
    if maxBetti0 == 0: maxBetti0=1
    if maxBetti1 == 0: maxBetti1=1
    if maxBetti2 == 0: maxBetti2=1
    if maxBetti3 == 0: maxBetti3=1
    if maxBetti4 == 0: maxBetti4=1
    if maxHoles == 0: maxHoles=1
    if maxVol == 0: maxVol=1
    n2 = 0
    while n2 < n1:
        nBetti0[n2] = nBetti0[n2] / maxBetti0
        nBetti1[n2] = nBetti1[n2] / maxBetti1
        nBetti2[n2] = nBetti2[n2] / maxBetti2
        nBetti3[n2] = nBetti3[n2] / maxBetti3
        nBetti4[n2] = nBetti4[n2] / maxBetti4
        nHoles[n2] = nHoles[n2] / maxHoles
        if volNorm == True: vol[n2] = vol[n2] / maxVol
        n2 += 1

    fig = plt.figure(figsize=(13,8))
    plt.gca().xaxis.set_major_formatter(mdates.DateFormatter('%Y/%m/%d'))
    qtInt = int(qtDias / 10)
    print(f"qtInt: {qtInt}")
    plt.gca().xaxis.set_major_locator(mdates.DayLocator(interval=qtInt))
    #plt.plot(xd, vol, 'b-', linewidth=1.5, alpha=0.7)

    graphText = f"N.Emp: {ba.nEmp}, Comp: {ba.lenInt}, Passo: {ba.stInt} ({diaI1} a {diaF1}) [edgeMax: {ba.edgeMax}, persLim: {ba.persLim}]"
    if "V" in graphs:
        if volNorm:
            plt.plot(xDatas, vol, color='red', label=f"Volume (F: {maxVol:.2f})")
        else:
            plt.plot(xDatas, vol, color='red', label='Volume')
    if "B0" in graphs: plt.plot(xDatas, nBetti0, color='black', label=f"Betti 0 (F: {maxBetti0})")
    if "B1" in graphs: plt.plot(xDatas, nBetti1, color='green', label=f"Betti 1 (F: {maxBetti1})")
    if "B2" in graphs: plt.scatter(xDatas, nBetti2, color='magenta', label=f"Betti 2 (F: {maxBetti2})")
    if "B3" in graphs: plt.scatter(xDatas, nBetti3, color='cyan', label=f"Betti 3 (F: {maxBetti3})")
    if "B4" in graphs: plt.scatter(xDatas, nBetti3, color='yellow', label=f"Betti 4 (F: {maxBetti4})")
    if "H" in graphs: plt.plot(xDatas, nHoles, color='blue', label=f"Buracos (F: {maxHoles})")

    plt.xlabel(f"Datas", fontsize=10)
    plt.gcf().autofmt_xdate()

    plt.title(graphText)
    plt.legend(loc='upper right', bbox_to_anchor=(1.11, 1.07), labelcolor='linecolor', fontsize=11, frameon=True, edgecolor='black', facecolor='lightgray')
    plt.grid(True, alpha=0.3)
    plt.show()

    #return xd

def smeJsonSave (obj, fname, indent=-1):
    import json

    if indent < 0:
        jsonStr = json.dumps(obj.__dict__)
    else:
        jsonStr = json.dumps(obj.__dict__, indent=indent)

    if fname == '': return jsonStr

    try:
        with open(fname, 'w', encoding="utf-8") as file:
            file.write(jsonStr)
    except:
        print(f"***** Erro a criar o ficheiro '{fname}'")
            
    return 1

def smeJsonLoad (fname):
    import json

    try:
        with open(fname, 'rb') as file:
            #jsonStr = file.read().replace('\n', '')
            jsonStr = file.read()
    except:
        print(f"***** Erro a ler o ficheiro '{fname}'")

    #d1 = json.JSONDecoder()
    #obj = d1.decode(jsonStr)
    obj = json.loads(jsonStr)
    return obj


def smeCsvSave (rows, fields, fname=''):
    import csv

    #rows = [{"id": 1, "name": "Alice"}, {"id": 2, "name": "Bob"}]
    #fields = ["id", "name"]

    try:
        with open("users.csv", "w", newline="", encoding="utf-8") as file:
            writer = csv.DictWriter(fname, fieldnames=fields)
            writer.writeheader()
            writer.writerows(rows)
    except:
        print("***** Erro a criar ficheiro")

    return
