#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
vs_batch.py —— 批量虚拟筛选模板（AutoDock Vina 命令行 + 多进程 + 结果汇总）

用途
    给一个已经准备好的受体（PDBQT）和一个装满配体（PDBQT）的文件夹，
    用多进程并行跑 Vina，把每个配体的最优打分汇总成一张 CSV。

依赖
    - AutoDock Vina 可执行文件（vina 或 vina.exe），并已加入 PATH
      安装：conda install -c conda-forge vina  或  pip install -U vina
    - 仅需 Python 标准库（subprocess / multiprocessing / csv / argparse / re）

典型用法
    # 1) 先准备受体（含极性氢）与配体（PDBQT）
    python vs_batch.py \
        --receptor 1iep_receptor.pdbqt \
        --ligands  ligands/ \
        --config   1iep_receptor.box.txt \
        --out-dir  poses \
        --csv      results.csv \
        --nproc    8 \
        --exhaustiveness 16 \
        --top-n    5

    # 2) 只想重新汇总已有结果，不重跑对接
    python vs_batch.py --ligands ligands/ --out-dir poses --csv results.csv --summarize-only

config 文件（--config 指向的文本文件）内容示例：
    center_x = 15.190
    center_y = 53.903
    center_z = 16.917
    size_x   = 20.0
    size_y   = 20.0
    size_z   = 20.0

若没有 config 文件，也可以直接用命令行给盒子：
    --center 15.190 53.903 16.917 --size 20 20 20
"""

import argparse
import csv
import os
import re
import shutil
import subprocess
import sys
from collections import Counter
from concurrent.futures import ProcessPoolExecutor, as_completed
from glob import glob

# 匹配输出 PDBQT 里的结果行，例如：
# REMARK VINA RESULT:    -13.2      0.000      0.000
_RESULT_RE = re.compile(
    r"^REMARK\s+VINA RESULT:\s+(-?\d+\.?\d*)\s+(-?\d+\.?\d*)\s+(-?\d+\.?\d*)"
)


def parse_scores(pdbqt_path):
    """读取一个 Vina 输出 PDBQT，返回 [(affinity, rmsd_lb, rmsd_ub), ...]（按出现顺序=按打分排序）。"""
    scores = []
    try:
        with open(pdbqt_path, "r", errors="ignore") as fh:
            for line in fh:
                m = _RESULT_RE.match(line.strip())
                if m:
                    scores.append((float(m.group(1)), float(m.group(2)), float(m.group(3))))
    except OSError:
        return []
    return scores


def find_ligands(paths):
    """把 --ligands 传进来的目录 / glob / 文件统一展开成 PDBQT 文件绝对路径列表。"""
    files = []
    for p in paths:
        if os.path.isdir(p):
            files.extend(glob(os.path.join(p, "*.pdbqt")))
        elif any(ch in p for ch in "*?["):
            files.extend(glob(p))
        elif os.path.isfile(p):
            files.append(p)
    # 去重并排序，保证可复现
    files = sorted(set(os.path.abspath(f) for f in files))
    # 排除已经带 _out 的结果文件，避免重复对接
    return [f for f in files if not os.path.basename(f).endswith("_out.pdbqt")]


def assign_stems(ligands):
    """给每个配体分配唯一的输出名（stem），避免不同目录下的同名配体互相覆盖。

    基线是文件名去掉扩展名；同一个 basename 出现多次时，统一追加 __1、__2 … 后缀。
    返回 (stems, dups)：
        stems —— {配体绝对路径: 输出名}
        dups  —— 出现冲突的 basename 列表（排序后），供提示用
    """
    bases = [os.path.splitext(os.path.basename(f))[0] for f in ligands]
    total = Counter(bases)
    seen = Counter()
    stems = {}
    for f, b in zip(ligands, bases):
        if total[b] == 1:
            stems[f] = b
        else:
            seen[b] += 1
            stems[f] = "%s__%d" % (b, seen[b])
    dups = sorted(b for b, c in total.items() if c > 1)
    return stems, dups


def build_box_config(box_path, center, size):
    """没有现成 config 时，按 center/size 生成一个临时 config 文件。"""
    with open(box_path, "w") as fh:
        fh.write("center_x = %.3f\n" % center[0])
        fh.write("center_y = %.3f\n" % center[1])
        fh.write("center_z = %.3f\n" % center[2])
        fh.write("size_x = %.1f\n" % size[0])
        fh.write("size_y = %.1f\n" % size[1])
        fh.write("size_z = %.1f\n" % size[2])
    return box_path


def run_one(job):
    """单个配体对接。job 是一个 dict，方便进程池序列化。"""
    ligand = job["ligand"]
    name = job.get("stem") or os.path.splitext(os.path.basename(ligand))[0]
    out_file = os.path.join(job["out_dir"], name + "_out.pdbqt")
    log_file = os.path.join(job["out_dir"], name + ".log")

    cmd = [
        job["vina"],
        "--receptor", job["receptor"],
        "--ligand", ligand,
        "--config", job["config"],
        "--out", out_file,
        "--exhaustiveness", str(job["exhaustiveness"]),
        "--cpu", str(job["cpu_per_job"]),
    ]
    if job["scoring"]:
        cmd += ["--scoring", job["scoring"]]
    if job["seed"] is not None:
        cmd += ["--seed", str(job["seed"])]

    try:
        with open(log_file, "w") as log:
            proc = subprocess.run(
                cmd, stdout=log, stderr=subprocess.STDOUT,
                timeout=job["timeout"], check=False,
            )
    except subprocess.TimeoutExpired:
        return {"ligand": name, "status": "timeout", "out_file": out_file}
    except FileNotFoundError:
        return {"ligand": name, "status": "vina-not-found", "out_file": out_file}

    if proc.returncode != 0 or not os.path.isfile(out_file):
        return {"ligand": name, "status": "failed(rc=%d)" % proc.returncode, "out_file": out_file}

    scores = parse_scores(out_file)
    if not scores:
        return {"ligand": name, "status": "no-result", "out_file": out_file}

    return {
        "ligand": name,
        "status": "ok",
        "out_file": out_file,
        "best_affinity": scores[0][0],
        "best_rmsd_lb": scores[0][1],
        "best_rmsd_ub": scores[0][2],
        "n_poses": len(scores),
    }


def summarize(out_dir, csv_path, top_n):
    """扫描 out_dir 下所有 *_out.pdbqt，汇总最优打分到 CSV。"""
    rows = []
    for out_file in sorted(glob(os.path.join(out_dir, "*_out.pdbqt"))):
        name = os.path.basename(out_file)[:-len("_out.pdbqt")]
        scores = parse_scores(out_file)
        if scores:
            best = scores[0]
            rows.append({
                "ligand": name,
                "best_affinity": best[0],
                "best_rmsd_lb": best[1],
                "best_rmsd_ub": best[2],
                "n_poses": len(scores),
                "out_file": out_file,
            })
    rows.sort(key=lambda r: r["best_affinity"])  # 越负越靠前

    with open(csv_path, "w", newline="") as fh:
        writer = csv.writer(fh)
        writer.writerow(["rank", "ligand", "best_affinity_kcal_mol",
                         "rmsd_lb", "rmsd_ub", "n_poses", "out_file"])
        for i, r in enumerate(rows, 1):
            writer.writerow([i, r["ligand"], r["best_affinity"],
                             r["best_rmsd_lb"], r["best_rmsd_ub"],
                             r["n_poses"], r["out_file"]])
    return rows


def main():
    ap = argparse.ArgumentParser(description="AutoDock Vina 批量虚拟筛选模板")
    ap.add_argument("--receptor", help="受体 PDBQT 文件")
    ap.add_argument("--ligands", nargs="+", required=True,
                    help="配体 PDBQT：目录 / glob / 文件，可给多个")
    ap.add_argument("--config", help="Vina config 文本（含 center/size）")
    ap.add_argument("--center", nargs=3, type=float, metavar=("X", "Y", "Z"),
                    help="盒子中心；与 --size 搭配，替代 --config")
    ap.add_argument("--size", nargs=3, type=float, metavar=("X", "Y", "Z"),
                    help="盒子边长（埃）；与 --center 搭配")
    ap.add_argument("--out-dir", default="poses", help="输出目录（默认 poses）")
    ap.add_argument("--csv", default="vs_results.csv", help="汇总 CSV 路径")
    ap.add_argument("--nproc", type=int, default=os.cpu_count() or 1,
                    help="并行进程数（同时跑多少个配体）")
    ap.add_argument("--cpu-per-job", type=int, default=1,
                    help="每个 vina 进程内部使用几个 CPU 核（默认 1）")
    ap.add_argument("--exhaustiveness", type=int, default=16,
                    help="搜索强度；越大越慢越稳（默认 16，正式筛选可用 32）")
    ap.add_argument("--scoring", default=None, choices=["vina", "ad4", "vinardo"],
                    help="打分函数；默认用 vina。用 ad4 时需先算 affinity maps 并改用 --maps")
    ap.add_argument("--seed", type=int, default=None, help="随机种子，固定可复现")
    ap.add_argument("--timeout", type=int, default=3600, help="单个配体超时秒数")
    ap.add_argument("--top-n", type=int, default=0, help="打印前 N 名（0=不打印）")
    ap.add_argument("--summarize-only", action="store_true",
                    help="跳过对接，只重新汇总 out-dir 里已有结果")
    ap.add_argument("--vina", default=None, help="vina 可执行文件路径（默认从 PATH 找）")
    args = ap.parse_args()

    os.makedirs(args.out_dir, exist_ok=True)
    vina_bin = args.vina or shutil.which("vina") or "vina"

    if args.summarize_only:
        rows = summarize(args.out_dir, args.csv, args.top_n)
        print("汇总完成：%d 条结果 -> %s" % (len(rows), args.csv))
        _print_top(rows, args.top_n)
        return

    if not args.receptor:
        ap.error("需要 --receptor（受体 PDBQT）")
    if not args.config:
        if not (args.center and args.size):
            ap.error("需要 --config，或同时给出 --center 与 --size")
        args.config = build_box_config(os.path.join(args.out_dir, "_box.txt"),
                                       args.center, args.size)

    ligands = find_ligands(args.ligands)
    if not ligands:
        print("没有找到任何配体 PDBQT，检查 --ligands 路径。", file=sys.stderr)
        sys.exit(1)

    # 提前校验输入：这些问题逐个配体报一遍既费时间又难定位，直接拦在前面
    if not os.path.isfile(args.receptor):
        ap.error("找不到受体文件：%s" % args.receptor)
    if not os.path.isfile(args.config):
        ap.error("找不到 config 文件：%s" % args.config)
    if not (os.path.isfile(vina_bin) or shutil.which(vina_bin)):
        ap.error("找不到 vina 可执行文件（%s）。请先 conda install -c conda-forge vina，"
                 "或用 --vina 指定完整路径。" % vina_bin)

    # 不同目录下的同名配体输出会互相覆盖，这里统一改名规避
    stems, dups = assign_stems(ligands)
    if dups:
        print("提示：以下配体重名，输出名已自动加 __N 后缀避免覆盖：%s"
              % ", ".join(dups))

    print("vina        : %s" % vina_bin)
    print("受体        : %s" % os.path.abspath(args.receptor))
    print("config      : %s" % os.path.abspath(args.config))
    print("配体数量    : %d" % len(ligands))
    print("并行进程    : %d（每进程 %d 核）" % (args.nproc, args.cpu_per_job))
    print("搜索强度    : %d" % args.exhaustiveness)

    jobs = [{
        "ligand": lig,
        "stem": stems[lig],
        "receptor": os.path.abspath(args.receptor),
        "config": os.path.abspath(args.config),
        "out_dir": os.path.abspath(args.out_dir),
        "vina": vina_bin,
        "exhaustiveness": args.exhaustiveness,
        "cpu_per_job": args.cpu_per_job,
        "scoring": args.scoring,
        "seed": args.seed,
        "timeout": args.timeout,
    } for lig in ligands]

    done = 0
    results = []
    with ProcessPoolExecutor(max_workers=args.nproc) as pool:
        futures = {pool.submit(run_one, j): j for j in jobs}
        for fut in as_completed(futures):
            r = fut.result()
            results.append(r)
            done += 1
            tag = "%s(%s)" % (r["ligand"], r["status"])
            print("[%d/%d] %s" % (done, len(jobs), tag))

    rows = summarize(args.out_dir, args.csv, args.top_n)
    ok = sum(1 for r in results if r["status"] == "ok")
    print("\n对接结束：成功 %d / 共 %d，汇总 -> %s" % (ok, len(jobs), args.csv))
    _print_top(rows, args.top_n)


def _print_top(rows, top_n):
    if top_n and rows:
        print("\n打分最优的前 %d 个配体（kcal/mol，越负越可能结合）：" % top_n)
        for i, r in enumerate(rows[:top_n], 1):
            print("  %2d. %-30s %8.2f" % (i, r["ligand"], r["best_affinity"]))


if __name__ == "__main__":
    main()
