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

import os
import numpy as np

from smeUtils import *
from smeGudhi import *

class AnalisadorBuracos:
    def __init__(self, pontos, dimensao_maxima=2):
        self.pontos = np.array(pontos)
        self.dimensao_maxima = dimensao_maxima
        self.rips_complex = None
        self.persistence = None
        self.barcode = None

    # Mostra um conjunto de pontos original
    def visualizar_pontos(self, titulo="Conjunto de Pontos"):
        plt.figure(figsize=(10, 8))
        
        if self.pontos.shape[1] == 2:
            plt.scatter(self.pontos[:, 0], self.pontos[:, 1], 
                       c='blue', s=100, alpha=0.7, edgecolors='black')
            plt.xlabel('X')
            plt.ylabel('Y')
        elif self.pontos.shape[1] == 3:
            ax = plt.subplot(111, projection='3d')
            ax.scatter(self.pontos[:, 0], self.pontos[:, 1], self.pontos[:, 2],
                      c='blue', s=100, alpha=0.7)
            ax.set_xlabel('X')
            ax.set_ylabel('Y')
            ax.set_zlabel('Z')
        else:
            # Para dimensões maiores, usa PCA
            from sklearn.decomposition import PCA
            pca = PCA(n_components=2)
            pontos_2d = pca.fit_transform(self.pontos)
            plt.scatter(pontos_2d[:, 0], pontos_2d[:, 1], 
                       c='blue', s=100, alpha=0.7, edgecolors='black')
            plt.xlabel(f'PC1 ({pca.explained_variance_ratio_[0]*100:.1f}%)')
            plt.ylabel(f'PC2 ({pca.explained_variance_ratio_[1]*100:.1f}%)')
        
        plt.title(titulo)
        plt.grid(True, alpha=0.3)
        plt.tight_layout()
        plt.show()

    def _sparse_auto(self, raio):
        """Calcula automaticamente o parâmetro sparse."""
        n_pontos = len(self.pontos)
        if n_pontos > 100:
            return 0.5  # Mais esparso para conjuntos grandes
        return 1.0  # Mais denso para conjuntos pequenos

    # Constrói o complexo de Vietoris-Rips com diferentes raios
    def construir_rips_complex(self, raio_maximo=None, passo=0.01):
        # Calcula raio máximo se não especificado
        if raio_maximo is None:
            distancias = pdist(self.pontos)
            raio_maximo = np.max(distancias) / 2
        
        # Cria o complexo Rips
        self.rips_complex = gudhi.RipsComplex(
            points=self.pontos,
            max_edge_length=raio_maximo * 2,  # O complexo usa diâmetros
            sparse=self._sparse_auto(raio_maximo)
        )
        
        # Constrói o complexo simplicial
        simplex_tree = self.rips_complex.create_simplex_tree(
            max_dimension=self.dimensao_maxima
        )
        
        return simplex_tree

    # Calcula a persistência dos buracos.
    def calcular_persistencia(self, simplex_tree=None):
        if simplex_tree is None:
            simplex_tree = self.construir_rips_complex()
        
        # Calcula a persistência (birth, death, dimension)
        self.persistence = simplex_tree.persistence()
        return self.persistence
        
        # Calcula barcode
        self.barcode = simplex_tree.persistence_intervals_in_dimension(0)  # Componentes conexas
        for d in range(1, self.dimensao_maxima + 1):
            dim_barcode = simplex_tree.persistence_intervals_in_dimension(d)
            self.barcode = np.vstack([self.barcode, dim_barcode]) if len(self.barcode) > 0 else dim_barcode
        
        return self.persistence

    # Mostra o diagrama de persistência
    def visualizar_diag_persistencia (self, raio_otimo=None):
        if self.persistence is None:
            print("Erro: Calcule a persistência primeiro!")
            return
        
        plt.figure(figsize=(12, 5))

        dims = {}
        for (dim, (birth, death)) in self.persistence:
            if dim not in dims:
                dims[dim] = []
            dims[dim].append((birth, death))
        
        # Plota cada dimensão
        cores = ['blue', 'red', 'green', 'purple', 'orange']
        subplot_idx = 1
        
        for dim in sorted(dims.keys()):
            plt.subplot(1, len(dims), subplot_idx)
            
            for birth, death in dims[dim]:
                if np.isinf(death):
                    # Componente conexa infinita
                    plt.plot([birth], [dim], 'o', color=cores[dim % len(cores)], markersize=10, label=f'Dim {dim}')
                else:
                    plt.plot([birth, death], [dim, dim], '-', color=cores[dim % len(cores)], linewidth=2)
                    plt.plot([birth, death], [dim, dim], 'o', color=cores[dim % len(cores)], markersize=6)
            
            if raio_otimo:
                plt.axvline(x=raio_otimo, color='red', linestyle='--', label=f'Raio = {raio_otimo:.3f}')
            
            plt.xlabel('Raio')
            plt.ylabel('Dimensão')
            plt.title(f'Dimensão {dim}')
            plt.grid(True, alpha=0.3)
            subplot_idx += 1
        
        plt.tight_layout()
        plt.show()

    # Encontra buracos significativos baseados num limiar de persistência
    def encontrar_buracos_significativos (self, limiar_persistencia=0.1):
        if self.persistence is None:
            print("Erro: Calcule a persistência primeiro!")
            return []
        
        buracos_significativos = []
        
        #for birth, death, dim in self.persistence:
        for (dim, (birth, death)) in self.persistence:
            if dim > 0 and not np.isinf(death):  # Ignora componentes conexas
                persistencia = death - birth
                if persistencia > limiar_persistencia:
                    pt_central = birth + (death - birth) / 2
                    buracos_significativos.append({'dimensão': dim, 'nascimento': birth, 'morte': death,
                                                   'persistência': persistencia, 'ponto_central': pt_central})
        
        # Ordena por persistência
        buracos_significativos.sort(key=lambda x: x['persistência'], reverse=True)
        
        return buracos_significativos


def geraCirc (qtPts, centro=[0,0], raio=1.5, ruido=0.2):
    import numpy as np

    angs = np.linspace(0, 2*np.pi,qtPts, endpoint=False)
    pts = np.column_stack([
        centro[0] + raio * np.cos(angs) + np.random.normal(0, ruido, n_pontos),
        centro[1] + raio * np.sin(angs) + np.random.normal(0, ruido, n_pontos)
    ])

    return pts

    
def gera2circs (qtPts, centros=[[-2,0],[2,0]], raio=1.5, ruido=0.2):
    import numpy as np
    pts = []
    for c1 in centros:
        v1 = []
        for p in range(0, qtPts):
            ang = 2 * np.pi * np.random.random()
            x = c1[0] + raio * np.cos(ang) + np.random.random() * ruido
            y = c1[1] + raio * np.sin(ang) + np.random.random() * ruido
            v1.append([x,y])
        len1 = len(v1)
        pts.append(v1)

    return pts


def geraEsfera (qtPts, raio=2, centro=[0,0,0], ruido=0.1):
    n1 = 0
    v1 = []
    while n1 < qtPts:
    #for p in range(0, 10):
        theta = np.pi * np.random.random()
        phi = 2 * np.pi * np.random.random()
        
        x = centro[0] + raio * np.sin(theta) * np.cos(phi) + np.random.random() * ruido
        y = centro[1] + raio * np.sin(theta) * np.sin(phi) + np.random.random() * ruido
        z = centro[2] + raio * np.cos(theta) + np.random.random() * ruido
        v1.append([x,y,z])
        n1 += 1
    return v1


def geraToro (qtPts, R=4, r=1, centro=[0,0,0], ruido=0.1):
    import numpy as np
    
    v1 = []
    n1 = 0
    while n1 < qtPts:
        u = 2 * np.pi * np.random.random()
        v = 2 * np.pi * np.random.random()
        x = (R + r * np.cos(v)) * np.cos(u) + np.random.random() * ruido
        y = (R + r * np.cos(v)) * np.sin(u) + np.random.random() * ruido
        z = r * np.sin(v) + np.random.random() * ruido
        v1.append([x,y,z])
        n1 += 1
    return v1
    

#################################################

def norma3 (v1, centro=0):
    d = 0
    if centro == 0:
        for x in v1:
            d += x * x
    else:
        lenV = len(v1)
        if lenV != len(centro):
            print("\n  ***** Dimensões incorrectas do 'centro'\n")
            return -1
        n1 = 0
        while n1 < lenV:
            d += (v1[n1] - centro[n1]) * (v1[n1] - centro[n1])
            n1 += 1
    d = np.sqrt(d)
    return d


def gera2esferas (qtPts=300, centros=[[-2,0,0],[2,0,0]], raio=1.5, ruido=0.2):
    import numpy as np
    
    v1 = esfera(qtPts=qtPts, raio=raio, centro=centros[0], ruido=ruido)
    v2 = esfera(qtPts=qtPts, raio=raio, centro=centros[1], ruido=ruido)
    v3 = v1 + v3
    
    return v3

def circ (qtPts=10, raio=2):
    import numpy as np
    n1 = 0
    v1 = []
    while n1 < qtPts:
        theta = 2 * np.pi * np.random.random()
        
        x = raio * np.cos(theta)
        y = raio * np.sin(theta)
        d2 = x * x + y * y
        #print(f"raio: {raio} ; theta: {theta:12.8f} ; cos(theta): {np.cos(theta):12.8f} ; sin(theta): {np.sin(theta):11.8f} ; x: {x:11.8f} ; y: {y:11.8f} ; d2: {d2:11.8f}")
        v1.append([x,y])

        #teste
        d2 = x * x + y * y
        d1 = np.sqrt(d2)
        #print (f"----> n1: {n1:2} ; d2: {d2:12.8f} ; d1: {d1:12.8f}  ==?==  [x,y,z]: [{x:11.8f},{y:11.8f}]  ;;;  theta: {theta:11.8f}")
        n1 += 1
        
    return v1

def xxx (medge=0.4):
    pontos = gera2circulos(centros=[[-3,0],[3,0]])
    print("Len pontos: ", len(pontos))
    pontos2 = np.vstack(pontos)
    print("Len pontos2: ", len(pontos2))
    smeGhudiPlotRipsComplexFromPts(pontos2, max_edge_length=medge, title="xx")
    plt.show()
    return pontos2

def xxx2 (raio=3):
    import math
    #raio = 2
    qtPts=10
    v1 = gera2esferas(qtPts=qtPts, centros=[[0,0,0],[0,0,0]], ruido=0.0, raio=raio)
    n1 = 0
    while n1 < qtPts:
        #d1 = norma3(v1[0][n1], centro=[0,0,0])

        d1 = np.sqrt(v1[n1][0] * v1[n1][0] + v1[n1][1] * v1[n1][1] + v1[n1][2] * v1[n1][2])
        
        #if math.isclose(d1, raio) != True:
        print(f"d(v1[0][{n1}]) = {d1} != raio: {raio}")
        n1 += 1

    return v1


def readDataFileNY (cols):
    fname = "NY_19890420__20090420__319_dim250.txt"
    data = smeFileReadNum(fname, cols=cols)
    return data
