#!/usr/bin/python3
"""Exercise the installed console script and every bundled prediction model."""

import csv
import io
import itertools
import os
from pathlib import Path
import subprocess
import tempfile

os.environ['OMP_NUM_THREADS'] = '1'
os.environ['OPENBLAS_NUM_THREADS'] = '1'
os.environ.pop('PYTHONPATH', None)
sequence = 'VSGLEQLESIINFEKLTEWTSSNV'
header = ['position', 'residue', 'cleav_prob', 'cleaved', 'protein_id']


def run(*args):
    return subprocess.check_output(['pepsickle', *args], text=True)


def check(text, protein_id):
    reader = csv.DictReader(io.StringIO(text), delimiter='\t')
    assert reader.fieldnames == header, reader.fieldnames
    rows = list(reader)
    assert len(rows) == len(sequence), len(rows)
    for position, row in enumerate(rows, 1):
        assert int(row['position']) == position
        assert row['residue'] == sequence[position - 1]
        assert row['protein_id'] == protein_id
        probability = float(row['cleav_prob'])
        assert 0 <= probability <= 1
        # Printed probabilities are rounded to four decimal places.
        if abs(probability - 0.25) > 0.0001:
            assert (row['cleaved'] == 'True') == (probability > 0.25)
    assert float(rows[-1]['cleav_prob']) == 0
    assert rows[-1]['cleaved'] == 'False'


assert '--model-type' in run('--help')
with tempfile.TemporaryDirectory() as directory:
    fasta = Path(directory) / 'input.fasta'
    output = Path(directory) / 'output.tsv'
    fasta.write_text(f'>example\n{sequence}\n', encoding='ascii')
    for model, proteasome, human_only in itertools.product(
            ('epitope', 'in-vitro', 'in-vitro-2'), ('C', 'I'), (False, True)):
        args = ['-m', model, '-p', proteasome, '-t', '0.25']
        if human_only:
            args.append('--human-only')
        check(run('-s', sequence, *args), 'None')
        assert run('-f', str(fasta), '-o', str(output), *args) == ''
        check(output.read_text(encoding='utf-8'), 'example')
        print(f'{model}, {proteasome}, human_only={human_only}: OK')
    for threshold in ('-0.1', '1.1', 'nan', 'abc'):
        result = subprocess.run(
            ['pepsickle', '-s', sequence, '-t', threshold],
            stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True,
        )
        assert result.returncode != 0, threshold
