#!/usr/bin/env python3
"""
compare_aln.py — сравнение двух множественных выравниваний (аналог VerAlign)

Использование:
    python compare_aln.py aln1.fasta aln2.fasta [label1] [label2]
    python compare_aln.py -h

Входные данные:
    Два FASTA-файла с выравниваниями одних и тех же последовательностей.
    Поддерживаемые заголовки:
      - UniProt: >sp|ACC|ID_SPECIES ...
      - ID со диапазоном: >ACNA_ECOLI/1-891 ...
      - PDBeFold: >PDB:1x03:A ...
      - прочие: >1x03_A_Name ...
    В каждом файле длины выровненных последовательностей должны совпадать.
    Набор последовательностей (по short_id) должен пересекаться.

Выходные данные (stdout):
    - Длины выравниваний
    - Число и % совпадающих колонок
    - Список блоков (s1,f1)=(s2,f2) с длиной >= 2
    - Одиночные совпадающие колонки вне блоков
    - Несовпадающие участки в первом выравнивании

Опционально (если указан -o FILE):
    В FILE записывается список пар (i, j) одинаково выровненных колонок.
"""

import argparse
import sys


def short_id_from_header(line):
    """Из строки >header извлекает короткий идентификатор последовательности."""
    raw = line[1:].strip()
    header = raw.split()[0]
    parts = raw.split('|')
    if len(parts) >= 3 and parts[0] in ('sp', 'tr', 'SP', 'TR'):
        # UniProt: sp|ACC|ID_SPECIES
        return parts[2].split()[0].split('/')[0]
    if header.upper().startswith('PDB:'):
        # PDBeFold: PDB:1x03:A
        return header.split(':')[1].lower()
    # >ACNA_ECOLI/1-891 или >ACNA_ECOLI
    return header.split('/')[0]


def parse_fasta(filepath):
    """
    Читает FASTA файл с выравниванием.
    Возвращает dict {short_id: gapped_sequence}.
    """
    seqs = {}
    current = None
    with open(filepath, encoding='utf-8') as f:
        for line in f:
            line = line.rstrip()
            if not line:
                continue
            if line.startswith('>'):
                current = short_id_from_header(line)
                seqs[current] = ''
            elif current is not None:
                seqs[current] += line.replace(' ', '')
    return seqs


def compare_alignments(aln1, aln2, label1='Aln1', label2='Aln2'):
    """
    Сравнивает два выравнивания одних и тех же последовательностей.

    Алгоритм:
      Для каждой колонки c1 в aln1 находим остатки всех последовательностей.
      Смотрим, в какую колонку c2 в aln2 попадают те же остатки.
      Если для всех последовательностей это одна и та же колонка c2 —
      пара (c1, c2) считается совпадающей.

    Возвращает: (blocks, singles, matching)
      blocks  — список (s1, f1, s2, f2, length) блоков длиной >= 2
      singles — список (c1, c2) одиночных совпадений вне блоков
      matching — полный список совпадающих пар (c1, c2), 1-based
    """
    seqs = sorted(set(aln1.keys()) & set(aln2.keys()))
    if not seqs:
        print('ОШИБКА: нет общих последовательностей между файлами.')
        print(f'  {label1}: {sorted(aln1.keys())}')
        print(f'  {label2}: {sorted(aln2.keys())}')
        return [], [], []

    # Проверка одинаковой длины внутри каждого файла
    for lab, aln in ((label1, aln1), (label2, aln2)):
        lengths = {s: len(aln[s]) for s in seqs if s in aln}
        if len(set(lengths.values())) != 1:
            print(f'ОШИБКА: в {lab} длины последовательностей различаются: {lengths}')
            return [], [], []

    ncols1 = len(aln1[seqs[0]])
    ncols2 = len(aln2[seqs[0]])

    # остаток -> колонка в aln2
    res_to_col2 = {}
    for s in seqs:
        res_to_col2[s] = {}
        ri = 0
        for col, aa in enumerate(aln2[s]):
            if aa != '-':
                res_to_col2[s][ri] = col
                ri += 1

    # колонка -> индекс остатка в aln1 / aln2
    col_to_res1 = {}
    col_to_res2 = {}
    for s in seqs:
        col_to_res1[s] = []
        ri = 0
        for aa in aln1[s]:
            col_to_res1[s].append(None if aa == '-' else ri)
            if aa != '-':
                ri += 1
        col_to_res2[s] = []
        ri = 0
        for aa in aln2[s]:
            col_to_res2[s].append(None if aa == '-' else ri)
            if aa != '-':
                ri += 1

    matching = []
    for c1 in range(ncols1):
        assignments = {s: col_to_res1[s][c1] for s in seqs
                       if col_to_res1[s][c1] is not None}
        if not assignments:
            continue
        c2_candidates = set()
        for s, ri in assignments.items():
            c2_candidates.add(res_to_col2[s].get(ri, None))
        if len(c2_candidates) != 1 or None in c2_candidates:
            continue
        c2 = list(c2_candidates)[0]
        assignments2 = {s: col_to_res2[s][c2] for s in seqs
                        if col_to_res2[s][c2] is not None}
        if assignments2 == assignments:
            matching.append((c1 + 1, c2 + 1))

    blocks = []
    if matching:
        bs1, bs2 = matching[0]
        ps1, ps2 = matching[0]
        for c1, c2 in matching[1:]:
            if c1 == ps1 + 1 and c2 == ps2 + 1:
                ps1, ps2 = c1, c2
            else:
                if ps1 - bs1 + 1 >= 2:
                    blocks.append((bs1, ps1, bs2, ps2, ps1 - bs1 + 1))
                bs1, bs2, ps1, ps2 = c1, c2, c1, c2
        if ps1 - bs1 + 1 >= 2:
            blocks.append((bs1, ps1, bs2, ps2, ps1 - bs1 + 1))

    in_block = set()
    for b in blocks:
        for i in range(b[0], b[1] + 1):
            in_block.add(i)
    singles = [(c1, c2) for c1, c2 in matching if c1 not in in_block]

    print(f"\n{'='*60}")
    print(f"Сравнение: {label1} vs {label2}")
    print(f"Общие последовательности ({len(seqs)}): {', '.join(seqs)}")
    print(f"Длина {label1}: {ncols1}  |  Длина {label2}: {ncols2}")
    print(f"Совпадающих колонок: {len(matching)}")
    print(f"  % от {label1}: {100 * len(matching) / ncols1:.1f}%")
    print(f"  % от {label2}: {100 * len(matching) / ncols2:.1f}%")

    print(f"\nБлоки (длина >= 2), по убыванию длины:")
    if blocks:
        print(f"  {'(s1,f1)':>14} = {'(s2,f2)':>14}  длина")
        for b in sorted(blocks, key=lambda x: -x[4]):
            print(f"  ({b[0]},{b[1]}) = ({b[2]},{b[3]})  {b[4]}")
    else:
        print("  блоков нет")

    print(f"\nВсего блоков: {len(blocks)}")
    print(f"Одиночных совпадений вне блоков: {len(singles)}")
    if singles:
        print("  " + ", ".join(f"({c1},{c2})" for c1, c2 in singles))

    blocks_sorted = sorted(blocks, key=lambda x: x[0])
    print(f"\nНесовпадающие участки в {label1}:")
    prev = 0
    for b in blocks_sorted:
        if b[0] > prev + 1:
            print(f"  {prev + 1}-{b[0] - 1} (длина {b[0] - 1 - prev})")
        prev = b[1]
    if prev < ncols1:
        print(f"  {prev + 1}-{ncols1} (длина {ncols1 - prev})")

    return blocks, singles, matching


def main():
    parser = argparse.ArgumentParser(
        description='Сравнение двух MSA одних и тех же последовательностей.',
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog=__doc__,
    )
    parser.add_argument('aln1', nargs='?', help='Первое выравнивание (FASTA)')
    parser.add_argument('aln2', nargs='?', help='Второе выравнивание (FASTA)')
    parser.add_argument('label1', nargs='?', default=None, help='Метка первого выравнивания')
    parser.add_argument('label2', nargs='?', default=None, help='Метка второго выравнивания')
    parser.add_argument(
        '-o', '--output',
        help='Файл для списка пар (i, j) одинаково выровненных колонок',
    )
    args = parser.parse_args()

    if not args.aln1 or not args.aln2:
        parser.print_help()
        sys.exit(1)

    lab1 = args.label1 or args.aln1
    lab2 = args.label2 or args.aln2
    A = parse_fasta(args.aln1)
    B = parse_fasta(args.aln2)
    blocks, singles, matching = compare_alignments(A, B, label1=lab1, label2=lab2)

    if args.output:
        with open(args.output, 'w', encoding='utf-8') as out:
            out.write('# matching columns (i, j) 1-based\n')
            for i, j in matching:
                out.write(f'{i}\t{j}\n')
        print(f'\nСписок пар (i, j) записан в {args.output}')


if __name__ == '__main__':
    main()
