#!/usr/bin/python
# -*- coding: utf-8 -*-
#author : ruixia niu

"""
RNA结合蛋白（RBP）结合基序发现流程
功能：通过比较"处理组（富集）"和"对照组（总RNA）"的5-mer频率，
计算富集倍数（R值），并迭代识别高置信度结合基序，生成sequence logo。
"""

import sys
import os
import pandas as pd
import numpy as np
from collections import Counter
from scipy import stats

# =============================================================================
# 输入参数解析
# =============================================================================
script, total, rich = sys.argv
# total: 对照组/总RNA的序列文件（FASTQ格式）
# rich: 处理组/富集（如RIP-seq）的序列文件（FASTQ格式）


# =============================================================================
# 函数1: 统计k-mer频率（归一化）
# =============================================================================
def stat_kmer(seqs, k):
    """
    统计一组序列中所有k-mer的出现频率（归一化为概率）
    参数:
        seqs: 序列列表（仅包含碱基，不含ID和质量行）
        k: k-mer长度（此处固定为5）
    返回:
        dict: {k-mer: 频率}
    """
    kmers = []
    for seq in seqs:
        # 只处理长度恰好为30nt的序列（适配RBP实验的读段长度）
        if len(seq) == 30:
            # 滑动窗口提取k-mer
            for j in range(0, len(seq) - k + 1):
                # 跳过含N的k-mer，确保序列质量
                if "N" not in seq[j:j+k]:
                    kmers.append(seq[j:j+k])
    
    # 使用pandas计算归一化频率（更高效）
    se = pd.Series(kmers)
    kmers_dict = dict(se.value_counts(normalize=True))
    return kmers_dict


# =============================================================================
# 函数2: 序列比对，判断两个k-mer的相对位移
# =============================================================================
def align(str1, str2):
    """
    判断短序列str2相对于str1的位移类型（用于基序扩展）
    参数:
        str1: 参考序列（当前最优k-mer）
        str2: 候选序列
    返回:
        str: 位移类型 ("misn1", "misn2_l", "misn2_r", "misn3_l", "misn3_r") 或 "NO"
             其中 misn2_l 表示在参考左侧错配1个碱基，类似地 r 表示右侧
    """
    # 计算无位移时的错配数（理论上str1 == str2时最优）
    misn1 = sum(1 for i in range(len(str1)) if str1[i] != str2[i])
    
    # 计算str2相对于str1左移1位时的错配数（str2中多出的碱基在右侧）
    misn2_l = sum(1 for i in range(len(str1)-1) if str1[i] != str2[i+1])
    
    # 计算str2相对于str1右移1位时的错配数（str2中多出的碱基在左侧）
    misn2_r = sum(1 for i in range(1, len(str1)) if str1[i] != str2[i-1])
    
    # 类似地，计算左移2位和右移2位的情况
    misn3_l = sum(1 for i in range(len(str1)-2) if str1[i] != str2[i+2])
    misn3_r = sum(1 for i in range(2, len(str1)) if str1[i] != str2[i-2])

    # 根据错配模式推断最优比对方式（此处逻辑保留原算法）
    if misn1 == 1:
        # 当正好有1个错配时，认为是单碱基变异
        return "misn1"
    if misn1 == 2:
        # 当有2个错配时，可能对应插入/删除（indel）
        if misn2_l == 1:
            return "misn2_l"
        elif misn2_r == 1:
            return "misn2_r"
        elif misn3_l == 0:
            return "misn3_l"
        elif misn3_r == 0:
            return "misn3_r"
    
    return "NO"


# =============================================================================
# 函数3: 从序列中移除指定的k-mer（用于迭代屏蔽）
# =============================================================================
def mask_kmer(seqs, kmer):
    """
    从一组序列中移除指定的k-mer（用于去除已发现的基序，以便发现次级基序）
    参数:
        seqs: 序列列表
        kmer: 要移除的k-mer字符串
    返回:
        list: 移除k-mer后的序列列表
    """
    rmseqs = []
    for seq in seqs:
        # 如果kmer在序列中，则循环删除所有出现位置
        if kmer in seq:
            while kmer in seq:
                idx = seq.index(kmer)
                # 删除该kmer（连接前后片段）
                seq = seq[:idx] + seq[idx+len(kmer):]
            # 删除所有kmer后，将剩余序列加入结果
            rmseqs.append(seq)
        else:
            # 如果序列中不包含该kmer，则原样保留
            rmseqs.append(seq)
    return rmseqs


# =============================================================================
# 函数4: 计算富集倍数（R值）和Z-score
# =============================================================================
def calculate_R(total, rich):
    """
    计算每个5-mer在富集组相对于总RNA组的富集倍数R = rich_freq / total_freq
    并计算R值的Z-score，用于排序和筛选。
    参数:
        total: 总RNA组的k-mer频率字典
        rich: 富集组的k-mer频率字典
    返回:
        pd.DataFrame: 包含 kmer, R值, R_Z_score 的表格（按Z-score降序排列）
    """
    all_kmers = set(total.keys()) & set(rich.keys())  # 取交集，确保在两个库中均出现
    
    kmer_R = {}
    for kmer in all_kmers:
        # R = 富集频率 / 总频率，反映了该k-mer在结合组分中的富集程度
        kmer_R[kmer] = rich[kmer] / total[kmer]
    
    # 计算所有R值的Z-score（标准化），用于评估显著性
    R_values = list(kmer_R.values())
    R_z_scores = stats.zscore(R_values)
    
    R_data = pd.DataFrame({
        'kmer': list(kmer_R.keys()),
        'R': R_values,
        'R_Z_score': R_z_scores
    })
    # 按Z-score降序排列，越靠前的k-mer越可能是核心结合基序
    R_data = R_data.sort_values(by='R_Z_score', ascending=False)
    return R_data


# =============================================================================
# 函数5: 核心迭代流程 —— 发现基序 → 屏蔽 → 再发现
# =============================================================================
def find_logo(total_seq, rich_seq):
    """
    一次迭代：根据当前序列集合，发现最高富集的5-mer，并进行基序扩展
    参数:
        total_seq: 总RNA组的序列列表（仅含碱基）
        rich_seq: 富集组的序列列表（仅含碱基）
    返回:
        list: [logo1_seq, rich1_kmer, R_data]
               logo1_seq: 扩展后的基序及其权重（用于生成weblogo）
               rich1_kmer: 本次迭代找到的核心5-mer
               R_data: 完整的富集分析结果表
    """
    # 1. 统计5-mer频率
    kmer_total = stat_kmer(total_seq, 5)
    kmer_rich = stat_kmer(rich_seq, 5)
    
    # 2. 计算富集倍数R和Z-score
    R_data = calculate_R(kmer_total, kmer_rich)
    
    # 3. 取Z-score最高的5-mer作为核心基序
    rich1_kmer = R_data.iloc[0, 0]
    # 4. 筛选Z-score >= 3的显著基序作为候选集
    top_kmers = R_data[R_data["R_Z_score"] >= 3]["kmer"].tolist()
    
    # 5. 计算每个显著基序的权重：权重 = R - 1（表示相对背景的富集程度）
    logo_weight_dict = {
        kmer: R_data.loc[R_data["kmer"] == kmer, "R"].values[0] - 1
        for kmer in top_kmers
    }
    
    # 6. 基序扩展：基于序列比对，从核心基序向左右扩展
    kmer_9_dict = {}      # 存储扩展的k-mer
    logo_seq_dict = {}    # 存储最终用于logo的序列及其权重
    
    # 遍历所有显著基序，判断它们与核心基序的位移关系
    for kmer in top_kmers:
        align_pos = align(rich1_kmer, kmer)
        if align_pos != "NO":
            # 初始化该k-mer的扩展序列列表
            if kmer not in kmer_9_dict:
                kmer_9_dict[kmer] = []
            
            # 如果是单碱基错配，直接添加该k-mer自身
            if align_pos == "misn1":
                kmer_9_dict[kmer].append(kmer)
            
            # 根据比对位移方向，从原始富集序列中提取对应的5-mer
            for seq in rich_seq:
                if "N" not in seq and len(seq) == 30:
                    if kmer in seq:
                        p = seq.index(kmer)
                        # 根据比对类型，提取左移/右移对应的5-mer
                        if align_pos == "misn2_l":
                            s = seq[p+1:p+1+5]
                            if len(s) == 5:
                                kmer_9_dict[kmer].append(s)
                        elif align_pos == "misn2_r":
                            s = seq[p-1:p-1+5]
                            if len(s) == 5:
                                kmer_9_dict[kmer].append(s)
                        elif align_pos == "misn3_l":
                            s = seq[p+2:p+2+5]
                            if len(s) == 5:
                                kmer_9_dict[kmer].append(s)
                        elif align_pos == "misn3_r":
                            s = seq[p-2:p-2+5]
                            if len(s) == 5:
                                kmer_9_dict[kmer].append(s)
            
            # 统计扩展序列的频率，并乘以对应权重
            if kmer in kmer_9_dict:
                se = pd.Series(kmer_9_dict[kmer])
                kmer_9_freq = dict(se.value_counts(normalize=True))
                for ext_kmer, freq in kmer_9_freq.items():
                    logo_seq_dict[ext_kmer] = logo_weight_dict[kmer] * freq
    
    # 添加核心基序本身，权重为其Z-score（原代码这里可能想用Z-score，但实际取了Z-score-1）
    logo_seq_dict[rich1_kmer] = R_data.iloc[0, 2] - 1
    
    return [logo_seq_dict, rich1_kmer, R_data]


# =============================================================================
# 主程序：迭代执行三次，生成最终的weblogo
# =============================================================================

# 读取序列文件（跳过FASTQ标识行和质量行，只取碱基序列）
# 注：原代码先读取了序列行（index=3），然后又读取了index=1，这是有问题的。
# 此处保留原逻辑不变，实际应明确需要读取哪一行（通常第2行是碱基序列）。
total_seq = []
with open(total) as f1:
    total_seqs = f1.readlines()
for l in range(0, len(total_seqs)-3, 4):
    total_seq.append(total_seqs[l+1].strip("\n"))  # FASTQ格式第2行是序列

rich_seq = []
with open(rich) as f2:
    rich_seqs = f2.readlines()
for l in range(0, len(rich_seqs)-3, 4):
    rich_seq.append(rich_seqs[l+1].strip("\n"))

# 迭代3次，每次发现一个基序并屏蔽后继续发现下一个
r = find_logo(total_seq, rich_seq)
i = 1
while i <= 3:
    # 生成输出文件名
    prefix = f"{rich.split('/')[-1].split('_')[0]}_{total.split('/')[-1].split('_')[0]}"
    logoname = f"{prefix}_logo_{i}.eps"
    logoseq = f"{prefix}_logo_{i}.fa"
    logoR = f"{prefix}_logo_{i}.txt"
    
    # 写入FASTA文件（用于weblogo）
    # 注意：原代码将T替换为U（RNA序列），同时按权重重复序列以反映丰度
    with open(logoseq, "w") as f1:
        for m, weight in r[0].items():
            # 将权重放大100倍，转换为重复次数（用于weblogo的序列权重）
            repeat_count = int(weight * 100)
            for n in range(repeat_count):
                f1.write(f">{m}{n}\n{m.replace('T','U')}\n")
    
    # 调用外部工具weblogo生成序列标识图（sequence logo）
    os.system(
        f"weblogo -f {logoseq} -A rna -S 2 --errorbars NO "
        f"--fineprint weglogo --annotate 1,2,3,4,5 -o {logoname}"
    )
    
    # 保存本次迭代的R值分析结果
    with open(logoR, "w") as f2:
        r[2].to_csv(f2, sep="\t", index=False)
    
    # 屏蔽已发现的基序，为下一轮迭代准备序列
    total_seq = mask_kmer(total_seq, r[1])
    rich_seq = mask_kmer(rich_seq, r[1])
    
    # 重新运行发现流程，寻找下一个次级基序
    r = find_logo(total_seq, rich_seq)
    i += 1
