####
#### Generate class I supplementary tables - (i) all SARS-CoV-2 epitopes ranked by population coverage, and (ii) SARS-CoV-2 epitopes per protein that maximize cumulative population coverage
####

import os
import sys
import pandas as pd
import numpy as np

import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from matplotlib.backends.backend_pdf import PdfPages
from matplotlib.legend_handler import HandlerLine2D
from textwrap import wrap
from matplotlib.pyplot import cm

population = 'averaged' # averaged indicates using an aggregate of USA, API, and EUR population frequencies

## load binding scores
df_scores = pd.read_csv('SARS-CoV-2.top2pct_neonmhc1.unique_peptide_alleles.scored.csv') # input: table with source protein(s), peptide, allele, neon_pctrnk 
df_scores = df_scores[df_scores['neon_pctrnk'] <= 1] # can further filter by %rank
peptides = sorted(set(df_scores['peptide']))

## load allele frequencies
df = pd.read_table('collated_classi_allele_frequencies.txt') # input: table of class I allele frequencies
df['allele'] = df['allele'].apply(lambda x: 'HLA-' + x.replace('*',''))
df['averaged'] = df[['EUR', 'API', 'USA']].mean(axis=1)

## rank peptides by population coverage, quality of binding
def loadFreqs(df, population):
  allele2freq = dict(zip(df['allele'], df[population]))
  missing_alleles = [i for i in set(df_scores['allele']) if i not in allele2freq] # patch frequencies for alleles with predictions that aren't in allele frequency table
  for allele in missing_alleles:
    allele2freq[allele] = 0
  return allele2freq

def calcCoverage(alleles, population):
  loci = ['HLA-A', 'HLA-B', 'HLA-C']
  product = 1
  for locus in loci:
    locus_alleles = [i for i in alleles if locus in i]
    summed_freqs = sum([allele2freq[i] for i in locus_alleles])
    product = product * (1 - summed_freqs)**2
  coverage = 1 - product
  return coverage

allele2freq = loadFreqs(df, population)

list_ = []
for peptide in peptides:
  alleles = ';'.join(sorted(df_scores[df_scores['peptide'] == peptide]['allele']))
  allele_coverage = calcCoverage(alleles.split(';'), population)
  scores = list(df_scores[df_scores['peptide'] == peptide]['neon_pctrnk']) # note scores aren't in same order as alleles
  median_score = np.median(scores)
  mean_score = np.mean(scores)
  min_score = np.min(scores)
  list_.append([peptide, alleles, allele_coverage, scores, median_score, mean_score, min_score])

df_flat = pd.DataFrame(list_, columns=['peptide', 'alleles', 'allele_coverage', 'neon_pctrnks', 'median_neon_pctrnks', 'mean_neon_pctrnks', 'min_neon_pctrnks'])
df_flat.sort_values(by=['median_neon_pctrnks'], inplace=True) # first sort by median pct ranks to get best average scoring peptides - could also score by best ranking overall
df_flat['EUR_coverage'] = df_flat['alleles'].apply(lambda x: calcCoverage(x.split(';'), 'EUR'))
df_flat['API_coverage'] = df_flat['alleles'].apply(lambda x: calcCoverage(x.split(';'), 'API'))
df_flat['USA_coverage'] = df_flat['alleles'].apply(lambda x: calcCoverage(x.split(';'), 'USA'))
df_flat = df_flat[['peptide', 'protein', 'alleles', 'allele_coverage', 'USA_coverage', 'EUR_coverage', 'API_coverage', 'neon_pctrnks', 'median_neon_pctrnks', 'mean_neon_pctrnks', 'min_neon_pctrnks']]

df_flat.sort_values(by=['allele_coverage'], inplace=True, ascending=False) # rank peptides by allele coverage
df_flat['hg19_human_proteome_overlap'] = df_flat['peptide'].isin(omitted_peptides) # omitted_peptides is the list of peptides that coincide with the human proteome
supp_cols = ['peptide', 'protein', 'alleles', 'USA_coverage', 'EUR_coverage', 'API_coverage', 'hg19_human_proteome_overlap']
df_flat[supp_cols].to_csv('classi_ranked_by_coverage.csv', sep=',', index=False) # output: corresponds to supplemental table of all top-ranking peptides sorted by allele coverage

## rank peptides that maximize coverage
peptides = list(df_flat['peptide'])
peptide2alleles = dict(zip(df_flat['peptide'], df_flat['alleles']))

def rankTopByGene(df, protein):
  if protein != '':
    df_flat = df[df['protein'] == protein]
  else:
    df_flat = df
  df_flat = df_flat[~df_flat['hg19_human_proteome_overlap']] # filter out human proteome collisions
  top_peptide = df_flat.iloc[0]['peptide']
  banked_alleles = df_flat[df_flat['peptide'] == top_peptide]['alleles'].iloc[0].split(';')
  banked_peptides = [top_peptide]
  while len(banked_peptides) < 10:
    df_ = df_flat[~df_flat['peptide'].isin(banked_peptides)]
    df_['new_alleles'] = df_['alleles'].apply(lambda x: [i for i in x.split(';') if i not in banked_alleles])
    df_['new_coverage'] = df_['new_alleles'].apply(lambda x: 1 - np.prod([1-allele2freq[i] for i in x]))
    df_.sort_values(by=['new_coverage'], inplace=True, ascending=False)
    top_peptide = df_['peptide'].iloc[0]
    banked_peptides.append(top_peptide)
    banked_alleles = sorted(set(banked_alleles + df_flat[df_flat['peptide'] == top_peptide]['alleles'].iloc[0].split(';')))
  list_ = []
  covered = []
  for h, peptide in enumerate(banked_peptides):
    alleles = peptide2alleles[peptide].split(';')
    covered = [i for i in alleles if i not in covered] + covered
    data = [h+1, peptide, pep2orf[peptide], ';'.join(sorted(set(covered)))]
    for population_ in ['USA', 'EUR', 'API']:
      coverage = calcCoverage(covered, population_)
      data.append(coverage)
    list_.append(data)
  df_cumulative = pd.DataFrame(list_, columns=['rank', 'peptide', 'protein', 'cumulative_alleles', 'USA_cumulative_coverage', 'EUR_cumulative_coverage', 'API_cumulative_coverage'])
  df_cumulative.to_csv(protein.replace(' ', '_') + '.ranked_by_cumulative_coverage.csv', sep=',', index=False) # output: corresponds to supplemental table of 10 peptides that maximize overall population coverage at protein level

proteins_of_interest = ['ORF9b protein', 'membrane glycoprotein', 'nucleocapsid phosphoprotein', 'surface glycoprotein', '']
for protein in proteins_of_interest:
  rankTopByGene(df_flat, protein)
