#!/usr/bin/env python3

from math import ceil
from bio.seq.base import Fasta, Falist, write_fasta_o

def split_by_num_fa(infa,outfile,n,lth):
    i = 0
    j = 0
    myout=open(outfile+str(i)+".fasta",'w')
    for seq in infa:
        j += 1
        if j > n:
            i += 1
            myout.close()
            myout=open(outfile+str(i)+".fasta",'w')
            j = 1
        write_fasta_o(myout,seq,lth=lth)

# Split to files, each contains n sequences.
def split_by_num(infile,outfile,n,lth):
    fa = Fasta(infile)
    n = int(n)
    split_by_num_fa(fa,outfile,n,lth)

# Split to n files.
def split_by_part(infile,outfile,n,lth):
    fa = Falist(infile)
    m = ceil(len(fa)/int(n))
    split_by_num_fa(fa,outfile,m,lth)

def main(argv):
    import argparse

    parser = argparse.ArgumentParser(description="Split fasta file to muilt files.")
    parser.add_argument('infile',help="Fasta file to be splited.")
    parser.add_argument('outfile',help="output file prefix")
    parser.add_argument('num', help="Each file num/ output files(set --bypart)")
    parser.add_argument('--bypart',action='store_true')
    parser.add_argument('--lth',nargs='?',default=80,type=int,help="lth")
    args = parser.parse_args(argv[1:])

    if args.bypart:
        split_by_part(args.infile,args.outfile,args.num,args.lth)
    else:
        split_by_num(args.infile,args.outfile,args.num,args.lth)

if __name__ == '__main__':
    import sys
    main(sys.argv)
