import os
import sys
import pickle
import numpy as np
import pandas as pd
import scanpy as sc
import pycisTopic
import logging
import sklearn.preprocessing as sp
import pyranges as pr
from scipy import sparse
from pycisTopic.cistopic_class import CistopicObject
from pycisTopic.lda_models import run_cgs_models_mallet
from pycisTopic.topic_binarization import binarize_topics
from pycisTopic.utils import (
    region_names_to_coordinates,
    subset_list,
    get_position_index,
    non_zero_rows
)

import argparse

parser = argparse.ArgumentParser(description="Run pycisTopic analysis for SCENIC+")
parser.add_argument("--project_dir", type=str, required=True,
                    help="Base project directory")
parser.add_argument("--scatac_h5ad", type=str, required=True,
                    help="Path to scATAC.h5ad file (from step 0)")
parser.add_argument("--blacklist", type=str, required=True,
                    help="Path to ENCODE blacklist BED file")
parser.add_argument("--mallet_dir", type=str, required=True,
                    help="Path to Mallet installation directory (contains bin/mallet)")
parser.add_argument("--java_home", type=str, default="",
                    help="Path to JAVA_HOME (default: use system JAVA_HOME)")
parser.add_argument("--mallet_memory", type=int, default=256,
                    help="Mallet memory in GB (default: 256)")
parser.add_argument("--n_topics", type=int, default=40,
                    help="Number of LDA topics (default: 40)")
parser.add_argument("--n_iter", type=int, default=500,
                    help="Number of Mallet iterations (default: 500)")
parser.add_argument("--random_state", type=int, default=555,
                    help="Random seed (default: 555)")
parser.add_argument("--n_cpu", type=int, default=32,
                    help="Number of CPUs (default: 32)")

args = parser.parse_args()

if args.java_home:
    os.environ["JAVA_HOME"] = args.java_home
    os.environ["PATH"] = os.environ["JAVA_HOME"] + "/bin:" + os.environ["PATH"]

PROJECT_DIR  = args.project_dir
OUTPUT_DIR   = os.path.join(PROJECT_DIR, "pycisTopic")
TEMP_DIR     = os.path.join(PROJECT_DIR, "tmp")

SCATAC_H5AD_PATH = args.scatac_h5ad
BLACKLIST_PATH   = args.blacklist

MALLET_DIR    = args.mallet_dir
MALLET_PATH   = os.path.join(MALLET_DIR, "bin", "mallet")
MALLET_MEMORY_GB = args.mallet_memory

N_TOPICS      = args.n_topics
N_ITER        = args.n_iter
RANDOM_STATE  = args.random_state
N_CPU         = args.n_cpu

def create_cistopic_object_custom(
    fragment_matrix,
    cell_names=None,
    region_names=None,
    path_to_blacklist=None,
    min_frag=1,
    min_cell=1,
    is_acc=1,
    path_to_fragments=None,
    project="cisTopic",
    tag_cells=True,
    split_pattern="___",
):
    level = logging.INFO
    log_format = "%(asctime)s %(name)-12s %(levelname)-8s %(message)s"
    handlers = [logging.StreamHandler(stream=sys.stdout)]
    logging.basicConfig(level=level, format=log_format, handlers=handlers)
    log = logging.getLogger("cisTopic")

    # DataFrame → sparse matrix
    if isinstance(fragment_matrix, pd.DataFrame):
        log.info("Converting fragment matrix to sparse matrix")
        region_names = list(fragment_matrix.index)
        cell_names = list(fragment_matrix.columns.values)
        fragment_matrix = sparse.csr_matrix(fragment_matrix.to_numpy(), dtype=np.int32)

    if tag_cells:
        cell_names = [cell_names[x] + split_pattern + project for x in range(len(cell_names))]

    # 黑名单过滤
    if isinstance(path_to_blacklist, str):
        log.info("Removing blacklisted regions")
        regions = pr.PyRanges(region_names_to_coordinates(region_names))
        blacklist = pr.read_bed(path_to_blacklist)
        regions = regions.overlap(blacklist, invert=True)
        if ':' in region_names[0]:
            separator = ':'
        else:
            separator = '-'
        selected_regions = (
            regions.Chromosome.astype(str)
            + separator
            + regions.Start.astype(str)
            + "-"
            + regions.End.astype(str)
        ).to_list()
        index = get_position_index(selected_regions, region_names)
        fragment_matrix = fragment_matrix[index,]
        region_names = selected_regions

    log.info("Creating CistopicObject")
    binary_matrix = sp.binarize(fragment_matrix, threshold=is_acc - 1)
    selected_regions = non_zero_rows(binary_matrix)
    fragment_matrix = fragment_matrix[selected_regions,]
    binary_matrix = binary_matrix[selected_regions,]
    region_names = subset_list(region_names, selected_regions)

    cisTopic_nr_frag = np.array(fragment_matrix.sum(axis=0)).flatten()
    cisTopic_nr_acc  = np.array(binary_matrix.sum(axis=0)).flatten()

    cell_data = pd.DataFrame(
        [cisTopic_nr_frag, np.log10(cisTopic_nr_frag),
         cisTopic_nr_acc,  np.log10(cisTopic_nr_acc),
         [project] * len(cell_names)],
        columns=cell_names,
        index=["cisTopic_nr_frag", "cisTopic_log_nr_frag",
               "cisTopic_nr_acc",  "cisTopic_log_nr_acc",
               "sample_id"],
    ).transpose()

    if min_frag != 1:
        selected_cells = cell_data.cisTopic_nr_frag >= min_frag
        fragment_matrix = fragment_matrix[:, selected_cells]
        binary_matrix = binary_matrix[:, selected_cells]
        cell_data = cell_data.loc[selected_cells,]
        cell_names = cell_data.index.to_list()

    region_data = region_names_to_coordinates(region_names)
    region_data["Width"] = abs(region_data.End - region_data.Start).astype(np.int32)
    region_data["cisTopic_nr_frag"] = np.array(fragment_matrix.sum(axis=1)).flatten()
    region_data["cisTopic_log_nr_frag"] = np.log10(region_data["cisTopic_nr_frag"])
    region_data["cisTopic_nr_acc"]  = np.array(binary_matrix.sum(axis=1)).flatten()
    region_data["cisTopic_log_nr_acc"] = np.log10(region_data["cisTopic_nr_acc"])

    if min_cell != 1:
        selected_regions = region_data.cisTopic_nr_acc >= min_cell
        fragment_matrix = fragment_matrix[selected_regions, :]
        binary_matrix = binary_matrix[selected_regions, :]
        region_data = region_data.loc[selected_regions, :]
        region_names = region_data.index.to_list()

    if path_to_fragments is None:
        path_to_fragments = {}

    cistopic_obj = CistopicObject(
        fragment_matrix, binary_matrix, cell_names, region_names,
        cell_data, region_data, path_to_fragments, project,
    )
    log.info("Done!")
    return cistopic_obj

def check_mallet():
    if not os.path.exists(MALLET_PATH):
        print(f"Error: Mallet executable not found at {MALLET_PATH}")
        sys.exit(1)
    print(f"Mallet found at: {MALLET_PATH}")
    return MALLET_PATH

def main():
    for d in [OUTPUT_DIR, TEMP_DIR]:
        if not os.path.exists(d):
            os.makedirs(d)

    mallet_path = check_mallet()
    os.environ['MALLET_HOME'] = MALLET_DIR
    os.environ['MALLET_MEMORY'] = f"{MALLET_MEMORY_GB}g"
    print(f"Setting Mallet memory to: {os.environ['MALLET_MEMORY']}")

    print(f"Loading ATAC data from {SCATAC_H5AD_PATH}...")
    if not os.path.exists(SCATAC_H5AD_PATH):
        print(f"Error: {SCATAC_H5AD_PATH} not found. Please run step 0 first.")
        sys.exit(1)

    adata = sc.read_h5ad(SCATAC_H5AD_PATH)
    print(f"Loaded data shape: {adata.shape}")

    print("Creating CistopicObject...")
    count_matrix = adata.X.T     
    cell_names   = adata.obs_names
    region_names = adata.var_names

    cistopic_obj = create_cistopic_object_custom(
        fragment_matrix=count_matrix,
        cell_names=cell_names.tolist(),
        region_names=region_names.tolist(),
        project="SCENICplus_Analysis",
        tag_cells=False,
        min_frag=0,
        min_cell=0,
        path_to_blacklist=BLACKLIST_PATH,
    )

    if not adata.obs.empty:
        cistopic_obj.add_cell_data(adata.obs)

    print(f"CistopicObject: {len(cistopic_obj.cell_names)} cells, {len(cistopic_obj.region_names)} regions.")

    print("Running Mallet topic modeling...")
    models = run_cgs_models_mallet(
        cistopic_obj,
        n_topics=[N_TOPICS],
        n_cpu=N_CPU,
        n_iter=N_ITER,
        random_state=RANDOM_STATE,
        alpha=50,
        alpha_by_topic=True,
        eta=0.1,
        eta_by_topic=False,
        tmp_path=TEMP_DIR,
        mallet_path=mallet_path,
    )

    if not hasattr(cistopic_obj, 'LDA_models'):
        cistopic_obj.LDA_models = {}
    if isinstance(models, list):
        for model in models:
            cistopic_obj.add_LDA_model(model)
            cistopic_obj.LDA_models[model.n_topic] = model
    else:
        cistopic_obj.add_LDA_model(models)
        cistopic_obj.LDA_models[models.n_topic] = models

    model_path = os.path.join(OUTPUT_DIR, "cistopic_obj_with_models.pkl")
    with open(model_path, 'wb') as f:
        pickle.dump(cistopic_obj, f)
    print(f"Saved CistopicObject with models to {model_path}")

    selected_model_ntopics = N_TOPICS
    print(f"Binarizing topics for model with {selected_model_ntopics} topics...")
    cistopic_obj.selected_model = cistopic_obj.LDA_models[selected_model_ntopics]

    region_bin_topics = binarize_topics(cistopic_obj, method='otsu', plot=False)

    binarized_path = os.path.join(OUTPUT_DIR, "binarized_topics.pkl")
    with open(binarized_path, 'wb') as f:
        pickle.dump(region_bin_topics, f)
    print(f"Saved binarized topics to {binarized_path}")

    consensus_bed_path = os.path.join(OUTPUT_DIR, "consensus_regions.bed")
    print(f"Exporting all regions to {consensus_bed_path}...")
    regions_df = pd.DataFrame(
        [r.replace(':', '-').split('-') for r in cistopic_obj.region_names],
        columns=['Chrom', 'Start', 'End']
    )
    regions_df.to_csv(consensus_bed_path, sep='\t', header=False, index=False)
    print(f"Exported {len(regions_df)} regions.")

    print("\n[Step 1] pycisTopic analysis completed successfully.")

if __name__ == "__main__":
    main()
