# smeGudhi
#
# def smeGhudiPlotRipsComplex (rips_complex, pts, maxDim=2, title="RipsComplex", figsize=(10, 8)):
#     Mostra um RipsComplex em 2D ou 3D. A dimensão obtém-se a partir de "pts"
#     pts: Array numpy de pontos (n x 2)[2D] ou (n x 3)[3D]
#     rips_complex: RipsComplex a fazer gráfico
#     maxDim: Dimensão máxima dos complexos simpliciais
#     title: Título do gráfico
#     figsize: Tamanho da figura
#
# def smeGhudiPlotRipsComplexFromPts (pts, max_edge_length, maxDim=2, title="RipsComplex", figsize=(10, 8), show=False):
#     Mostra de um RipsComplex em 2D ou 3D a partir dos seus pontos. A dimensão obtém-se a partir de "pts"
#     pts: Array numpy de pontos (n x 2)[2D] ou (n x 3)[3D]
#     max_edge_length: Comprimento máximo das arestas
#     maxDim: Dimensão máxima dos complexos simpliciais
#     title: Título do gráfico
#     figsize: Tamanho da figura
#

import gudhi
import matplotlib.pyplot as plt
import numpy as np
from mpl_toolkits.mplot3d import Axes3D


def smeGhudiPlotRipsComplex (rips_complex, pts, maxDim=2, title="RipsComplex", figsize=(10, 8), show=False):
    # Mostra um RipsComplex em 2D ou 3D. A dimensão obtém-se a partir de "pts"
    # pts: Array numpy de pontos (n x 2)[2D] ou (n x 3)[3D]
    # rips_complex: RipsComplex a fazer gráfico
    # maxDim: Dimensão máxima dos complexos simpliciais
    # title: Título do gráfico
    # figsize: Tamanho da figura
    
    # Criar SimplexTree
    #rips_complex = gudhi.RipsComplex(points=pts, max_edge_length=max_edge_length)
    simplex_tree = rips_complex.create_simplex_tree(max_dimension=maxDim)
    
    # Determinar dimensão dos dados
    dim = pts.shape[1]
    
    fig = plt.figure(figsize=figsize)
    
    if dim == 2:
        ax = fig.add_subplot(111)
        
        # Plotar pontos
        ax.scatter(pts[:, 0], pts[:, 1], s=100, c='red', zorder=5)
        
        # Plotar arestas
        for simplex in simplex_tree.get_skeleton(1):
            if len(simplex[0]) == 2:
                edge = simplex[0]
                x = [pts[edge[0]][0], pts[edge[1]][0]]
                y = [pts[edge[0]][1], pts[edge[1]][1]]
                ax.plot(x, y, 'b-', linewidth=1.5, alpha=0.7)
        
        ax.set_xlabel('X')
        ax.set_ylabel('Y')
        
    else:  # 3D
        ax = fig.add_subplot(111, projection='3d')
        
        # Plotar pontos
        ax.scatter(pts[:, 0], pts[:, 1], pts[:, 2], 
                   s=100, c='red', zorder=5)
        
        # Plotar arestas
        for simplex in simplex_tree.get_skeleton(1):
            if len(simplex[0]) == 2:
                edge = simplex[0]
                x = [pts[edge[0]][0], pts[edge[1]][0]]
                y = [pts[edge[0]][1], pts[edge[1]][1]]
                z = [pts[edge[0]][2], pts[edge[1]][2]]
                ax.plot(x, y, z, 'b-', linewidth=1.5, alpha=0.7)
        
        ax.set_xlabel('X')
        ax.set_ylabel('Y')
        ax.set_zlabel('Z')
    
    ax.set_title(title)
    ax.grid(True, alpha=0.3)
    plt.tight_layout()
    
    if show: plt.show()
    
    return fig, ax


def smeGhudiPlotRipsComplexFromPts (pts, max_edge_length, maxDim=2, title="RipsComplex", figsize=(10, 8), show=False):
    # Mostra de um RipsComplex em 2D ou 3D a partir dos seus pontos. A dimensão obtém-se a partir de "pts"
    # pts: Array numpy de pontos (n x 2)[2D] ou (n x 3)[3D]
    # max_edge_length: Comprimento máximo das arestas
    # maxDim: Dimensão máxima dos complexos simpliciais
    # title: Título do gráfico
    # figsize: Tamanho da figura
    import gudhi
    import numpy as np
    
    # Criar RipsComplex e árvore de simplex associada
    rips_complex = gudhi.RipsComplex(points=pts, max_edge_length=max_edge_length)
    simplex_tree = rips_complex.create_simplex_tree(max_dimension=maxDim)

    if type(pts) != np.ndarray:
        pts = np.vstack(pts)
    
    # Determinar dimensão dos dados
    dim = pts.shape[1]
    
    fig = plt.figure(figsize=figsize)
    
    if dim == 2:
        ax = fig.add_subplot(111)
        
        # Plotar pontos
        ax.scatter(pts[:, 0], pts[:, 1], s=100, c='red', zorder=5)
        
        # Plotar arestas
        for simplex in simplex_tree.get_skeleton(1):
            if len(simplex[0]) == 2:
                edge = simplex[0]
                x = [pts[edge[0]][0], pts[edge[1]][0]]
                y = [pts[edge[0]][1], pts[edge[1]][1]]
                ax.plot(x, y, 'b-', linewidth=1.5, alpha=0.7)
        
        ax.set_xlabel('X')
        ax.set_ylabel('Y')
        
    else:  # 3D
        ax = fig.add_subplot(111, projection='3d')
        
        # Plotar pontos
        ax.scatter(pts[:, 0], pts[:, 1], pts[:, 2], 
                   s=100, c='red', zorder=5)
        
        # Plotar arestas
        for simplex in simplex_tree.get_skeleton(1):
            if len(simplex[0]) == 2:
                edge = simplex[0]
                x = [pts[edge[0]][0], pts[edge[1]][0]]
                y = [pts[edge[0]][1], pts[edge[1]][1]]
                z = [pts[edge[0]][2], pts[edge[1]][2]]
                ax.plot(x, y, z, 'b-', linewidth=1.5, alpha=0.7)
        
        ax.set_xlabel('X')
        ax.set_ylabel('Y')
        ax.set_zlabel('Z')
    
    ax.set_title(title)
    ax.grid(True, alpha=0.3)
    plt.tight_layout()

    if show: plt.show()
    
    return fig, ax
