import sys
import copy
from ged4py.parser import GedcomReader
import os.path
import scipy
import numpy as np

# Trace a gedcom file from one base person
# find generations backwards
# find shortest paths to other persons, given as arg
#
# by Kaj Holmberg
#
# Usage:
# python traceged1.py file.ged meth perslab1 [perslabl2,3,4...]
#   find path from perslab1 tp perslabl2,3,4... in file.ged
#        if perslab1 has extension txt, read from file
#   meth=0: onlyup, not marriage
#   meth=1: onlydown, not marriage
#   meth=2: upanddown, not marriage
#   meth=3: sideways, via marriage
#   meth=4: costs
#   meth=5: up from two persons, find common root

# only follow tree upwards, i.e. parents
onlyup=1
# only follow tree downwards, i.e. children
onlydown=1
upanddown=0
# go via marriages also
marr=0
usecosts=0
prino=0

pri=0
writenamelist=1
writedist1=1
writedist2=1
writedist3=1
writeunreach=1
allab=[]
allnames=[]
#allnum=[]
alldad=[]
allmom=[]
allbirth=[]
alldeath=[]
allfam=[]

def namefromix(ix):
    # get string with name from index ix
    name=allnames[ix]
    lab=allab[ix]
    str1=str(lab)+' '+str(name)
    return(str1)
    
def namefromlab(numlab):
    # get string with name from label numlab
    if numlab in allab:
        ixn=allab.index(numlab)
        name=allnames[ixn]
        str1=str(numlab)+' '+str(name)
    else:
        str1=''
    return(str1)
    
def getstrpersix(ix):
    # get printstring for person with index ix
    name=allnames[ix]
    if alldad[ix] in allab:
        dadix=allab.index(alldad[ix])
        dadname=allnames[dadix]
    else:
        dadix=''
        dadname='-'
    if allmom[ix] in allab:
        momix=allab.index(allmom[ix])
        momname=allnames[momix]
    else:
        momix=''
        momname='-'
    #str1='Name '+str(name)+' '+str(allab[ix])
    str1=str(name)+' '+str(allab[ix])
    str1=str1+' '+str(allbirth[ix])+' - '+str(alldeath[ix])
    str1=str1+' dad: '+str(dadname)+' '+str(alldad[ix])
    str1=str1+' mom: '+str(momname)+' '+str(allmom[ix])
    str1=str1+'\n'
    return(str1)

def printpersix(ix):
    # print person ix
    str1=getstrpersix(ix)
    print(str1)

def printname(name):
    # print person name
    if name in allnames:
        ix=allnames.index(name)
        printpersix(ix)
        
def writedistonfile1(futname0,ix0,ynp0,pred0):
    print('Write distances on file '+str(futname0))
    pers0=allab[ix0]
    fut=open(futname0,'w')
    #ix0=allab.index(pers0)
    str1='Direct distances from '+str(allnames[ix0])+' '+str(pers0)+':\n\n'
    fut.write(str1)
    for i in range(nn):
        str1=allab[i]+' '+str(allnames[i])
        str1=str1+', dad: '+alldad[i]+', mom: '+allmom[i]
        if pred0[i]>=0:
            str1=str1+', dist: '+str(int(ynp0[i]))+', pred: '+str(pred0[i])
        else:
            str1=str1+', dist: -, pred: -'
        str1=str1+'\n'
        fut.write(str1)
    fut.close()

def writeunreachedonfile(futname0,ix0,pred0):
    print('Write unreached on file '+str(futname0))
    pers0=allab[ix0]
    fut=open(futname0,'w')
    #ix0=allab.index(pers0)
    str1='Unreached from '+str(allnames[ix0])+' '+str(pers0)+':\n\n'
    fut.write(str1)
    for i in range(nn):
        if pred0[i]<0:
            if allab[i]!=perslab1:
                str1=allab[i]+' '+str(allnames[i])+'\n'
                fut.write(str1)
    fut.close()

def writedistonfile2(futname0,ix0,ynp0,pred0):
    print('Write distances on file '+str(futname0))
    pers0=allab[ix0]
    fut=open(futname0,'w')
    #ix0=allab.index(pers0)
    str1='Direct distances from '+str(allnames[ix0])+' '+str(pers0)+':\n '+str(noy)+' persons\n'
    fut.write(str1)
    for i in range(nn):
        str1=allab[i]+' '+str(allnames[i])
        str1=str1+', dad: '+alldad[i]+', mom: '+allmom[i]
        if pred0[i]>=0:
            str1=str1+', dist: '+str(int(ynp0[i]))+', pred: '+str(pred0[i])
            str1=str1+'\n'
            fut.write(str1)
    fut.close()

def writedistonfile3(futname0,ix0,ynp0,pred0):
    print('Write distances on file '+str(futname0))
    pers0=allab[ix0]
    fut=open(futname0,'w')
    #ix0=allab.index(pers0)
    str1='Direct distances from '+str(allnames[ix0])+' '+str(pers0)+':\n '+str(noy)+' persons\n'
    fut.write(str1)
    for i in range(nn):
        if alldad[i]!='':
            if alldad[i] in allab:
                ix1=allab.index(alldad[i])
                nix1=allnames[ix1]
            else:
                nix1='-'
        if allmom[i]!='':
            if allmom[i] in allab:
                ix2=allab.index(allmom[i])
                nix2=allnames[ix2]
            else:
                nix2='-'
        str1=allab[i]+' '+str(allnames[i])
        str1=str1+', dad: '+alldad[i]+' '+str(nix1)
        str1=str1+', mom: '+allmom[i]+' '+str(nix2)
        if pred0[i]>=0:
            str1=str1+', dist: '+str(int(ynp0[i]))+', pred: '+str(pred0[i])
        else:
            str1=str1+', dist: -, pred: -'
        str1=str1+'\n'
        fut.write(str1)
    fut.close()

def writegenonfile(futname0,ix0,ynp0,pred0):
    # write generations
    npmax=0
    for i in range(nn):
        if pred[i]>=0:
            if int(ynp[i])>npmax:
                npmax=int(ynp[i])
    print('Up to '+str(npmax)+' generations')
    #print(ynp0)
    print('Write generations on file '+str(futname0))
    fut=open(futname0,'w')
    for k in range(npmax+1):
        iii2=[]
        for i in range(nn):
            if ynp0[i]==k:
                iii2.append(i)
                
        ny2=len(iii2)
        #print(ny2,iii2)
        if ny2>0:
            str1=''
            for i in range(ny2):
                ii=iii2[i]
                str1=str1+' '+allnames[ii]+': '+str(int(ynp0[ii]))+'\n'
            fut.write('Distance '+str(k)+': ('+str(ny2)+')\n')
            fut.write(str1)
            if pri:
                print('Distance '+str(k)+': ('+str(ny2))
                print(str1)
    fut.close()
                
def writenopathonfile(futname0,ix1,ix2):
    strp1=namefromix(ix1)
    strp2=namefromix(ix2)
    print('No path found from '+strp1+' to '+strp2)
    fut=open(futname0,'w')
    fut.write('No path found from '+strp1+' to '+strp2+'\n')
    fut.close()

def writepathonfile(futname0,ix1,ix2,path0):
    strp1=namefromix(ix1)
    strp2=namefromix(ix2)
    pn=len(path0)
    print('Found a path from '+strp1+' to '+strp2+' with '+str(pn)+' steps')
    if pri:
        for i in range(pn-1,-1,-1):
            ix=path0[i]
            print(allab[ix],allnames[ix])

    #futname=fname1+'-'+str(perslab2)+'-path1.txt'
    print('Write path on file '+str(futname0))
    fut=open(futname0,'w')
    fut.write('Path from '+strp1+' to '+strp2+':\n')
    fut.write(str(pn)+' steps \n')
    previx=-1
    for i in range(pn-1,-1,-1):
        ix=path0[i]
        if previx>=0:
            # marr?
            ic=amat[previx][ix]+amat[ix][previx]
            if ic>2:
                fut.write('Married\n')
            #else:
            #    if previx>=0:
            #        if allbirth[previx]>allbirth[ix]:
            #            fut.write('Down\n')
            #        else:
            #            fut.write('Up\n')
        ii=pn-i
        str1=str(str(ii)+' '+allab[ix])+' '+str(allnames[ix])
        str1=str1+' '+str(allbirth[ix])+'-'+str(alldeath[ix])+'\n'
        fut.write(str1)
        previx=copy.copy(ix)
    fut.close()
    
def addpathonfile(futname0,ix1,ix2,path0):
    strp1=namefromix(ix1)
    strp2=namefromix(ix2)
    pn=len(path0)
    print('Found a path from '+strp1+' to '+strp2+' with '+str(pn)+' steps')
    if pri:
        for i in range(pn-1,-1,-1):
            ix=path0[i]
            print(allab[ix],allnames[ix])

    #futname=fname1+'-'+str(perslab2)+'-path1.txt'
    print('Write path on file '+str(futname0))
    fut=open(futname0,'a+')
    fut.write('-----\n')
    fut.write('Path from '+strp1+' to '+strp2+':\n')
    fut.write(str(pn)+' steps \n')
    previx=-1
    for i in range(pn-1,-1,-1):
        ix=path0[i]
        if previx>=0:
            # marr?
            ic=amat[previx][ix]
            if ic>1: fut.write('Married\n')
        ii=pn-i
        str1=str(str(ii)+' '+allab[ix])+' '+str(allnames[ix])
        str1=str1+' '+str(allbirth[ix])+'-'+str(alldeath[ix])+'\n'
        fut.write(str1)
        previx=copy.copy(ix)
    fut.close()
    

# --- main ---

# python traceged1.py file.ged meth perslab1 [perslabl2,3,4...]
#   find path from perslab1 tp perslabl2,3,4... in file.ged
#   meth=0: onlyup, not marriage
#   meth=1: onlydown, not marriage
#   meth=2: upanddown, not marriage
#   meth=3: sideways, via marriage
#   meth=4: costs
perslabl2=[]
np2=0
narg=len(sys.argv)
if narg>2:
    meth=int(sys.argv[2])
else:
    meth=0
if meth==0:
    onlyup=1
    onlydown=0
    upanddown=0
    marr=0
    usecosts=0
elif meth==1:
    onlyup=0
    onlydown=1
    upanddown=0
    marr=0
    usecosts=0
elif meth==2:
    onlyup=0
    onlydown=0
    upanddown=1
    marr=0
    usecosts=0
elif meth==3:
    onlyup=0
    onlydown=0
    upanddown=0
    marr=1
    usecosts=0
elif meth==4:
    onlyup=0
    onlydown=0
    upanddown=0
    marr=1
    usecosts=1
elif meth==5:
    onlyup=1
    onlydown=0
    upanddown=0
    marr=0
    usecosts=0
else:
    onlyup=0
    onlydown=0
    upanddown=0
    marr=0
    usecosts=0

if narg>3:
    perslab1=sys.argv[3]
else:
    perslab1=''

if narg>4:
    pers2a=sys.argv[4]
    if os.path.splitext(pers2a)[1]=='.txt':
        # read names from file
        if os.path.isfile(pers2a):
            fin=open(pers2a,'r')
            for line in fin:
                line2=line.split()
                num1=line2[-1]
                perslabl2.append(num1)
    else:
        for i in range(4,narg):
            perslabl2.append(sys.argv[i])
    np2=len(perslabl2)
print(str(np2)+' targets')

str1=' *** Find paths'
if onlyup:
    str1=str1+' up'
if onlydown:
    str1=str1+' down'
if upanddown:
    str1=str1+' up-down'
if marr:
    str1=str1+' incl marriage'
if usecosts:
    str1=str1+' with costs'
print(str1)

fname=os.path.basename(sys.argv[1])
print('Reading file '+str(fname))
fname=os.path.splitext(fname)[0]
fname1=fname+'+'+str(meth)+'-'+str(perslab1)

# read data
print('Reading individuals')
# open GEDCOM file
with GedcomReader(sys.argv[1]) as parser:
    # iterate over each INDI record in a file
    for i, indi in enumerate(parser.records0("INDI")):
        nra=indi.xref_id
        nr=nra.strip('@')
        name1=indi.name.format()
        if len(name1)>1:
            allab.append(nr)
            allnames.append(name1)
            #allnum.append(i)

            father = indi.father
            if father:
                dad=father.name.format()
                nra=father.xref_id
                nr=nra.strip('@')
            else:
                dad=''
                nr=''
            alldad.append(nr)
            
            mother = indi.mother
            if mother:
                mom=mother.name.format()
                nra=mother.xref_id
                nr=nra.strip('@')
            else:
                mom=''
                nr=''
            allmom.append(nr)

            birth = indi.sub_tag_value("BIRT/DATE")
            if birth:
                bir=f"{birth}"
            else:
                bir=''
            allbirth.append(bir)
            
            death = indi.sub_tag_value("DEAT/DATE")
            if death:
                dead=f"{death}"
            else:
                dead=''
            alldeath.append(dead)
            
    if marr:
        # get families
        print('Reading families')
        
        # iterate over each FAM record in a file
        for i, fam in enumerate(parser.records0("FAM")):

            husband, wife = fam.sub_tag("HUSB"), fam.sub_tag("WIFE")
            if husband: 
                hus=f"    husband: {husband.name.format()}"
                nra=husband.xref_id
                nr1=nra.strip('@')
            else:
                hus=''
                nr1=''
            if wife: 
                wif=f"    wife: {wife.name.format()}"
                nra=wife.xref_id
                nr2=nra.strip('@')
            else:
                wif=''
                nr2=''
                
            allfam.append([nr1,nr2])

        
nn=len(allnames)
print(str(nn)+' names read')
nf=len(allfam)
print(str(nf)+' families read')

# print names
if pri:
    for i in range(nn):
        printpersix(i)
    print('----------------------')
if writenamelist:
    futname=fname+'-names.txt'
    print('Write names on file '+str(futname))
    fut=open(futname,'w')
    for i in range(nn):
        str1=getstrpersix(i)
        fut.write(str1)
    fut.close()

# write node names
futname=fname+'-nodes.txt'
print('Write nodes on file '+str(futname))
fut=open(futname,'w')
fut.write(str(nn)+'\n')
for i in range(nn):
    str1=allab[i]+'\n'
    fut.write(str1)
fut.close()

# make links
print('Make link list')
startn=[]
endn=[]
costs=[]
for i in range(nn):
    ix0=allab.index(allab[i])
    if alldad[i]!='':
        if alldad[i] in allab:
            ix1=allab.index(alldad[i])
            if onlyup:
                startn.append(ix0)
                endn.append(ix1)
            elif onlydown:
                startn.append(ix1)
                endn.append(ix0)
            else:
                startn.append(ix0)
                endn.append(ix1)
            costs.append(1)
    if allmom[i]!='':
        if allmom[i] in allab:
            ix2=allab.index(allmom[i])
            #startn.append(ix0)
            #endn.append(ix2)
            if onlyup:
                startn.append(ix0)
                endn.append(ix2)
            elif onlydown:
                startn.append(ix2)
                endn.append(ix0)
            else:
                startn.append(ix0)
                endn.append(ix2)
            costs.append(1)
nl1=len(startn)
# families
if marr:
    for i in range(nf):
        hus,wif=allfam[i]
        if hus!='' and wif!='':
            if (hus in allab) and (wif in allab):
                # add link
                ix1=allab.index(hus)
                ix2=allab.index(wif)
                startn.append(ix1)
                endn.append(ix2)
                costs.append(10)
fut.close()
nl=len(startn)
print(str(nl)+' links created')

# make node matrix
print('Make node matrix')
amat=np.zeros((nn,nn),dtype=np.uint8)
for i in range(nl): 
    #amat[startn[i]][endn[i]]=1
    amat[startn[i]][endn[i]]=costs[i]

## increase cost for marriage
#if marr and costs:
#    for i in range(nl1,nl): 
#        amat[startn[i]][endn[i]]=10

    
# start at person 1
if perslab1=='':
    sys.exit()
    
if perslab1 in allab:
    pix1=allab.index(perslab1)
else:
    print('Person '+str(perslab1)+' not in list')
    sys.exit()

strp1=namefromlab(perslab1)
print('*** From person '+str(strp1))

print('Find shortest path')

from scipy.sparse import csr_matrix
a_sparse=csr_matrix(amat)

from scipy.sparse.csgraph import dijkstra
if meth==0 or meth==1 or meth==5:
    ynp,pred = scipy.sparse.csgraph.dijkstra(a_sparse, directed=True, indices=pix1, return_predecessors=True, unweighted=True)
elif meth==2:
    ynp,pred = scipy.sparse.csgraph.dijkstra(a_sparse, directed=False, indices=pix1, return_predecessors=True, unweighted=True)
elif meth==3:
    ynp,pred = scipy.sparse.csgraph.dijkstra(a_sparse, directed=False, indices=pix1, return_predecessors=True, unweighted=True)
elif meth==4:
    ynp,pred = scipy.sparse.csgraph.dijkstra(a_sparse, directed=False, indices=pix1, return_predecessors=True, unweighted=False)
if pri: print('Node prices:',ynp)
if pri: print('Pred:',pred)

# write result
if pri:
    for i in range(nn):
        if pred[i]>=0:
            print(allab[i],allnames[i],alldad[i],allmom[i],ynp[i],pred[i])

noy=0
for i in range(nn):
    if pred[i]>=0:
        noy+=1

#pix1=allab.index(perslab1)
if writedist1:
    writedistonfile1(fname1+'-dist1.txt',pix1,ynp,pred)
if writeunreach:
    writeunreachedonfile(fname1+'-unreach.txt',pix1,pred)
if writedist2:
    writedistonfile2(fname1+'-dist2.txt',pix1,ynp,pred)
if writedist3:
    writedistonfile3(fname1+'-dist3.txt',pix1,ynp,pred)

#ynpp=[]
#ixx=[]
#for i in range(nn):
#    if pred[i]>=0:
#        ynpp.append(int(ynp[i]))
#        ixx.append(i)
#nnpp=len(ynpp)
#if nnpp>0:
#    npmax=max(ynpp)
#else:
#    npmax=0

# generations
writegenonfile(fname1+'-gen1.txt',pix1,ynp,pred)


# person 2...

if np2==0:
    sys.exit()

for j in range(np2):
    perslab2=perslabl2[j]
    if perslab2 in allab:
        pix2=allab.index(perslab2)
        strp2=namefromlab(perslab2)

        if meth<5:
            # end at person 2
            #pix1=allab.index(perslab1)
            print('\n*** To person '+str(strp2))

            # find path backwards
            done=0
            path1=[]
            y1=[]
            fail=0
            ix1=copy.copy(pix2)
            path1.append(ix1)
            y1.append(ynp[ix1])
            while not done:
                ix2=pred[ix1]
                #print(ix1,ix2)
                if ix2==pix1:
                    # path ready
                    path1.append(ix2)
                    y1.append(ynp[ix2])
                    done=1
                elif ix2<0:
                    if pri: print('No path found')
                    fail=1
                    done=1
                else:
                    path1.append(ix2)
                    y1.append(ynp[ix2])
                    ix1=copy.copy(ix2)

            if fail:
                if prino:
                    futname=fname1+'-'+str(perslab2)+'-path1.txt'
                    writenopathonfile(futname,pix1,pix2)
            else:
                #print(y1)
                futname=fname1+'-'+str(perslab2)+'-path1.txt'
                writepathonfile(futname,pix1,pix2,path1)
        else:
            # meth 5
            #pix1=allab.index(perslab1)
            #strp2=namefromlab(perslab2)
            print('\n*** Now from person '+str(strp2))

            # first check if direct ancestor
            if pix2 in pred:
                print(strp2+' is direct ancestor to '+strp1)
                # ordinary path

                # find path backwards
                done=0
                path1=[]
                fail=0
                ix1=copy.copy(pix2)
                path1.append(ix1)
                while not done:
                    ix2=pred[ix1]
                    #print(ix1,ix2)
                    if ix2==pix1:
                        # path ready
                        path1.append(ix2)
                        done=1
                    elif ix2<0:
                        if pri: print('No path found')
                        fail=1
                        done=1
                    else:
                        path1.append(ix2)
                        ix1=copy.copy(ix2)

                if fail:
                    if prino:
                        futname=fname1+'-'+str(perslab2)+'-path1.txt'
                        writenopathonfile(futname,pix1,pix2)
                else:
                    futname=fname1+'-'+str(perslab2)+'-path1.txt'
                    writepathonfile(futname,pix1,pix2,path1)
            else:
                # not direct ancestor
                print('Find shortest path')
                ynp2,pred2 = scipy.sparse.csgraph.dijkstra(a_sparse, directed=True, indices=pix2, return_predecessors=True, unweighted=True)

                if pri: print('Node prices:',ynp2)
                if pri: print('Pred:',pred2)

                # write result
                if pri:
                    for i in range(nn):
                        if pred2[i]>=0:
                            print(allab[i],allnames[i],alldad[i],allmom[i],ynp2[i],pred2[i])
                        
                noy=0
                for i in range(nn):
                    if pred2[i]>=0:
                        noy+=1

                if writedist1:
                    writedistonfile1(fname1+'-dist12.txt',pix2,ynp2,pred2)
                if writeunreach:
                    writeunreachedonfile(fname1+'-unreach2.txt',pix2,pred2)
                if writedist2:
                    writedistonfile2(fname1+'-dist22.txt',pix2,ynp2,pred2)
                if writedist3:
                    writedistonfile3(fname1+'-dist32.txt',pix2,ynp2,pred2)


                #ynpp2=[]
                #ixx2=[]
                #for i in range(nn):
                #    if pred2[i]>=0:
                #        ynpp2.append(int(ynp2[i]))
                #        ixx2.append(i)
                #nnpp2=len(ynpp2)
                #if nnpp2>0:
                #    npmax2=max(ynpp2)
                #else:
                #    npmax2=0

                # generations
                writegenonfile(fname1+'-gen12.txt',pix2,ynp2,pred2)


                # now find common ancestor, person 3
                predinter=set(pred).intersection(set(pred2))
                #print(predinter)
                if len(predinter)==0:
                    print('No common ancestor')
                    pix3=-9999
                elif len(predinter)==1:
                    pix3=list(predinter)[0]
                else:
                    # more than one, find first
                    npi=len(predinter)
                    geni=999999
                    pix3=-1
                    for pe in predinter:
                        #print(pe)
                        if pe>-1:
                            if ynp[pe]<geni:
                                geni=ynp[pe]
                                #print(pe,ynp[pe],geni)
                                pix3=copy.copy(pe)
                            #elif ynp[pe]==geni:
                            #    if ynp2[pe]<geni:
                            #        geni=ynp[pe]
                            #        pix3=copy.copy(pe)
                if pix3<0:
                    print('No common ancestor')
                else:
                    strp3=namefromix(pix3)
                    print('Common ancestor '+strp3+' at gen '+str(geni))
                    # found ancestor
                
                    # path pers1 - pers 3

                    # find path backwards
                    done=0
                    path2=[]
                    fail=0
                    ix1=copy.copy(pix3)
                    path2.append(ix1)
                    while not done:
                        ix2=pred[ix1]
                        #print(ix1,ix2)
                        if ix2==pix1:
                            # path ready
                            path2.append(ix2)
                            done=1
                        elif ix2<0:
                            if pri: print('No path found')
                            fail=1
                            done=1
                        else:
                            path2.append(ix2)
                            ix1=copy.copy(ix2)

                    if fail:
                        if prino:
                            futname=fname1+'-'+str(perslab2)+'-path11.txt'
                            writenopathonfile(futname,pix1,pix3)
                    else:
                        #futname=fname1+'-'+str(perslab2)+'-path11.txt'
                        #writepathonfile(futname,pix1,pix3,path2)
                        futname=fname1+'-'+str(perslab2)+'-path13.txt'
                        writepathonfile(futname,pix1,pix3,path2)


                    # path pers2 - pers 3

                    # find path backwards
                    done=0
                    path12=[]
                    fail=0
                    ix1=copy.copy(pix3)
                    path12.append(ix1)
                    while not done:
                        ix2=pred2[ix1]
                        #print(ix1,ix2)
                        if ix2==pix2:
                            # path ready
                            path12.append(ix2)
                            done=1
                        elif ix2<0:
                            if pri: print('No path found')
                            fail=1
                            done=1
                        else:
                            path12.append(ix2)
                            ix1=copy.copy(ix2)

                    if fail:
                        if prino:
                            futname=fname1+'-'+str(perslab2)+'-path12.txt'
                            writenopathonfile(futname,pix2,pix3)
                    else:
                        #futname=fname1+'-'+str(perslab2)+'-path12.txt'
                        #writepathonfile(futname,pix2,pix3,path12)
                        futname=fname1+'-'+str(perslab2)+'-path13.txt'
                        addpathonfile(futname,pix2,pix3,path12)

    else:
        print('Person '+str(perslab2)+' not in list')

