"""Compare exactly two protein FASTA records with explicit identity denominators."""
import argparse
from itertools import islice
import json
from pathlib import Path
import sys

import Bio
from Bio import SeqIO
from Bio.Align import PairwiseAligner, substitution_matrices

ALPHABET = set("ACDEFGHIKLMNPQRSTVWY")


def read_pair(path, max_length):
    if max_length < 1:
        raise ValueError("max length must be positive")
    with Path(path).open(encoding="utf-8-sig") as handle:
        if handle.read(1) != ">":
            raise ValueError("FASTA must begin with a header")
        handle.seek(0)
        records = list(islice(SeqIO.parse(handle, "fasta"), 3))
    if len(records) != 2:
        raise ValueError("supply exactly two FASTA records: query, then target")
    if not all(record.id for record in records) or records[0].id == records[1].id:
        raise ValueError("record IDs must be nonempty and distinct")
    sequences = [str(record.seq).upper() for record in records]
    for record, sequence in zip(records, sequences):
        if not sequence or len(sequence) > max_length:
            raise ValueError(f"{record.id}: sequence length must be 1..{max_length}")
        invalid = set(sequence) - ALPHABET
        if invalid:
            raise ValueError(f"{record.id}: unsupported symbols {''.join(sorted(invalid))!r}")
    return records, sequences


def compare(path, mode="local", max_length=5000):
    records, (query, target) = read_pair(path, max_length)
    aligner = PairwiseAligner(mode=mode)
    aligner.substitution_matrix = substitution_matrices.load("BLOSUM62")
    aligner.open_gap_score = -10.0
    aligner.extend_gap_score = -0.5
    alignment = next(iter(aligner.align(target, query)), None)
    if alignment is None:
        raise ValueError("no positive-scoring local alignment")
    paired = matches = 0
    for (t0, t1), (q0, q1) in zip(*alignment.aligned):
        left, right = target[t0:t1], query[q0:q1]
        paired += len(left)
        matches += sum(a == b for a, b in zip(left, right))
    if paired == 0:
        raise ValueError("alignment has no residue-to-residue columns")
    columns = int(alignment.length)
    result = {
        "query_id": records[0].id, "target_id": records[1].id,
        "query_length": len(query), "target_length": len(target),
        "mode": mode, "matrix": "BLOSUM62", "gap_open": -10.0, "gap_extend": -0.5,
        "biopython": Bio.__version__, "score": float(alignment.score),
        "matches": matches, "paired_residues": paired, "alignment_columns": columns,
        "identity_paired_percent": round(100 * matches / paired, 2),
        "identity_columns_percent": round(100 * matches / columns, 2),
        "query_paired_coverage_percent": round(100 * paired / len(query), 2),
        "target_paired_coverage_percent": round(100 * paired / len(target), 2),
        "query_span": [int(x) for x in alignment.coordinates[1, [0, -1]]],
        "target_span": [int(x) for x in alignment.coordinates[0, [0, -1]]],
    }
    return result, str(alignment)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("fasta", type=Path, help="two records: query, then target")
    parser.add_argument("--mode", choices=["local", "global"], default="local")
    parser.add_argument("--max-length", type=int, default=5000)
    parser.add_argument("--show", action="store_true", help="print alignment to stderr")
    args = parser.parse_args()
    try:
        result, alignment = compare(args.fasta, args.mode, args.max_length)
    except (OSError, ValueError) as error:
        print(f"ERROR: {error}", file=sys.stderr)
        return 1
    if args.show:
        print(alignment, file=sys.stderr)
    print(json.dumps(result, indent=2))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
