import sys
import numpy as np
from mpi4py import MPI

if len(sys.argv)<2:
    print("Usage: filename")
    exit()
else:
    filename=sys.argv[1]

comm=MPI.COMM_WORLD
rank=comm.Get_rank()
nprocs=comm.Get_size()

root=0

my_char=chr(ord('A') + rank).encode('ascii')

# Set up the buffer to contain a one-byte character (consistent with C char)
buf=np.array([my_char])

status=MPI.Status()
#Info is used for "hints" to various MPI routines.  Usually we don't need it.
info=MPI.Info()

amode=MPI.MODE_CREATE | MPI.MODE_WRONLY
fh=MPI.File.Open(comm,filename,amode,info)

#Use MPI Byte type since Python doesn't really support characters
nreps=20
for i in range(nreps):
    offset=rank+i*nprocs
    fh.Write_at(offset, [buf, MPI.BYTE],status=status)

fh.Close()

if rank==0:
    #First read back with ordinary IO
    with open(filename,'rb') as fp:
        array=np.fromfile(fp,dtype='byte')
    array=[chr(array[i]) for i in range(array.size)]
    print("Read at root "+"".join(array))
    fp.close()

#All processes read entire file
#Blocking so will wait for root to finish above read
amode=MPI.MODE_RDONLY
fh=MPI.File.Open(comm,filename,amode)
fsize=fh.Get_size()
itembytes=MPI.BYTE.Get_size()
if rank==root:
    print("File size is "+str(fsize)+" type size is "+str(itembytes)+" bytes")
nbytes=fsize//itembytes
rbuf=np.empty((nbytes,),dtype='byte')
fh.Read_all([rbuf,MPI.BYTE])
all_chars=[chr(rbuf[i]) for i in range(rbuf.size)]
print("Entire file at rank "+str(rank)+" "+''.join(all_chars))
fh.Close()

#Read back the MPI file portion for each rank
amode=MPI.MODE_RDONLY
fh=MPI.File.Open(comm,filename,amode)

my_vals=[]
offset=0
for i in range(nreps):
    offset=rank+i*nprocs
    fh.Read_at(offset, [buf, MPI.BYTE])
    my_vals.append(buf[0].decode())
fh.Close()
my_string=''.join(my_vals)
print(str(rank)+' '+my_string)
