import networkx as nx

def dist3D(p, q):
    # Retorna la distància 3D entre p i q
    x1, y1, z1 = p
    x2, y2, z2 = q
    return ( (x1-x2) ** 2 + (y1-y2) ** 2 + (z1-z2) ** 2 ) ** 0.5


def crea_graf_complet(nomf):
    g = nx.Graph()
    # Llegim el fitxer i afegim un node al graf amb les coordenades de cada punt de consum
    with open(nomf, 'r') as f:
        for linia in f:
            linia = linia.split()
            coords = tuple(map(int, linia))
            g.add_node(coords)
    # Afegim al graf una aresta entre cada parell de nodes amb l'atribut de la distància
    for punt1 in g:
        for punt2 in g:
            if punt1 != punt2 and not g.has_edge(punt1, punt2):
                d = dist3D(punt1, punt2)
                g.add_edge(punt1, punt2, dist=d)
    return g

def cablejat(g):
    return nx.minimum_spanning_tree(g, weight='dist')
