#!/usr/bin/env python3

import random

def sample(lst,n):
    if n > len(lst):
        return lst
    else:
        return random.sample(lst,n)

def read_set(infile,has_header=False):
    data = set()
    if has_header:
        header = infile.readline()
    for line in infile:
        data.add(line.strip())
    return data

def parse_name(name):
    code = name.strip().split("_")[-1]
    assert len(code) == 8

    spe = code[:3]
    db = code[3]
    lth = code[4]
    tp = code[5:]
    # species, database, length, transcrpt type
    return spe, db, lth, tp

# 0->Ensmble 1->Noncode 3->NM 4->NR
def class_NM_NR_Ensb_Nonc(name):
    spe, db, lth, tp = parse_name(name)
    # Ensmble
    if db == '0':
        return '0'
    # Noncode
    elif db == '1':
        return '1'
    # refseq
    elif db == '2':
        # NM
        if tp == '001':
            return '3'
        # NR
        elif tp == '120':
            return '4'
        # XM or XR
        else:
            return None
    # db error
    else:
        raise ValueError("DB Error")

# outdata:{spe:{db:[name]}}
# db: Ensb, Nonc, NM, NR
def mk_sp_db(names,keep_sp):
    data={}
    for name in names:
        spe = parse_name(name)[0]
        if spe not in keep_sp:
            continue
        db = class_NM_NR_Ensb_Nonc(name)
        if db == None:
            continue
        dbdata = data.setdefault(spe,{'0':[],'1':[],'3':[],'4':[]})
        dbdata[db].append(name.strip())

    return data

def mk_sp_db1(names):
    data={}
    for name in names:
        spe = parse_name(name)[0]
        db = class_NM_NR_Ensb_Nonc(name)
        if db == None:
            continue
        dbdata = data.setdefault(spe,{'0':[],'1':[],'3':[],'4':[]})
        dbdata[db].append(name.strip())

    return data

# data:{spe:{db:[name]}}
def write_names(data,outfile,n=25):
    for val1 in data.values():
        for val2 in val1.values():
            if not val2:
                continue
            outfile.write("\n".join(sample(val2,n))+"\n")

def main(argv):
    import argparse

    parser = argparse.ArgumentParser(description="Split fasta file.")
    parser.add_argument('infa',nargs='?',help="fasta file for sampling ",default=sys.stdin,type=argparse.FileType('r'))
    parser.add_argument('-k','--keepspe',nargs='?',help="code file",type=argparse.FileType('r'))
    parser.add_argument('-o','--outfile',nargs='?',help="output file",default=sys.stdout,type=argparse.FileType('w'))
    parser.add_argument('-n','--num',nargs='?',help="how many sequence should be get",type=int,default=25)
    args = parser.parse_args(argv[1:])

    myfa = args.infa
    if args.keepspe:
        db = read_set(args.keepspe)
        data = mk_sp_db(myfa,db)
    else:
        data = mk_sp_db1(myfa)
    write_names(data,sys.stdout,args.num)

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