import pandas as pd
import os
import sys
import warnings
import json
import re
import hashlib
warnings.filterwarnings("ignore")
import logging
from spatialsnake.workflow.function.get_sample import (
    build_subset_analysis_tasks,
    get_sample_paths as _get_sample_paths,
    get_annotation,
    get_stereoseq_input_spec_map as _get_stereoseq_input_spec_map,
    read_compare_sample_table,
)
from spatialsnake.workflow.function.stereoseq_spec import parse_stereoseq_input_spec
L = logging.getLogger("spatialsnake_user")
L.setLevel(logging.INFO)
L.propagate = False
log_handler = logging.StreamHandler(sys.stdout)
formatter = logging.Formatter('%(asctime)s: %(levelname)s - %(message)s')
log_handler.setFormatter(formatter)
if not L.handlers:
    L.addHandler(log_handler)
runtype = ["visium","xenium","Merfish","visium_segment","stereoseq"]

annotation_list = config.get("annotation_list", "annotation.txt")
def parse_bool_flag(value):
    if isinstance(value, bool):
        return value
    if value is None:
        return False
    normalized = str(value).strip().lower()
    if normalized in {"true", "1", "yes", "y", "t"}:
        return True
    if normalized in {"false", "0", "no", "n", "f", "none", "null", ""}:
        return False
    return bool(value)

filter_list = parse_bool_flag(config.get("filter_list", False))
sample_list = config.get("sample_list", "sample.txt")
if not sample_list or not os.path.isfile(sample_list):
    sys.exit(
        "\n未提供必要样本信息，请提供 sample.txt 或通过命令指定样本表路径"
        f"；当前样本表路径不可用: {sample_list or 'sample.txt'}"
    )

spatialsnake_path = config.get("spatialsnake_path", "")
results_folder = config.get("results_folder", "results")
cellchat_compare_output_dir = config.get("cellchat_compare_output_dir", os.path.join(results_folder, "compare_cellchat"))
cellchat_compare_sample_name1 = config.get("cellchat_compare_sample_name1", "")
cellchat_compare_sample_name2 = config.get("cellchat_compare_sample_name2", "")
cellchat_compare_pathways = config.get("cellchat_compare_pathways", "")
cellchat_compare_focus_cells = config.get("cellchat_compare_focus_cells", "")
cellchat_compare_cell_pairs = config.get("cellchat_compare_cell_pairs", "")
cellchat_compare_source_cells = config.get("cellchat_compare_source_cells", "")
cellchat_compare_target_cells = config.get("cellchat_compare_target_cells", "")
cellchat_compare_lr_pairs = config.get("cellchat_compare_lr_pairs", "")
cellchat_compare_top_cell_pairs = config.get("cellchat_compare_top_cell_pairs", 3)
cellchat_compare_top_pathways = config.get("cellchat_compare_top_pathways", 3)
cellchat_compare_top_lr = config.get("cellchat_compare_top_lr", 20)
cellchat_compare_plot_advanced = config.get("cellchat_compare_plot_advanced", True)
cellchat_compare_receiver_cells = config.get("cellchat_compare_receiver_cells", "")
cellchat_compare_bubble_angle = config.get("cellchat_compare_bubble_angle", 45)
cellchat_compare_bubble_remove_isolate = config.get("cellchat_compare_bubble_remove_isolate", True)
cellchat_compare_do_single_bubble = config.get("cellchat_compare_do_single_bubble", False)
cellchat_compare_gene_colors = config.get("cellchat_compare_gene_colors", "white,#FEC44F,#D95F0E")
cellchat_compare_gene_plot_type = config.get("cellchat_compare_gene_plot_type", "dot")
cellchat_compare_pair_lr_use = config.get("cellchat_compare_pair_lr_use", "")
cellchat_compare_save_merged = config.get("cellchat_compare_save_merged", True)

option = config.get("option","integrate")

channel = config.get("channel", "single_analysis")

data_fold = config.get("data_fold", "data")
run_type = config.get("run_type", "visium")
#################  xenium ###########
cells_boundaries = config.get("cells_boundaries", False)
nucleus_boundaries = config.get("nucleus_boundaries", False)
nucleus_labels = config.get("nucleus_labels", False)
morphology_mip = config.get("morphology_mip", False)

markers_algorithm = config.get("markers_algorithm", "wilcoxon")
image_slice = config.get("image_slice", False)

if image_slice == "True":
    coord = []
    x1 = config.get("x1")
    x2 = config.get("x2")
    y1 = config.get("y1")
    y2 = config.get("y2")
    coord.append(x1)
    coord.append(x2)
    coord.append(y1)
    coord.append(y2)
tsne = config.get("tsne", False)
MIN_DIST = config.get("MIN_DIST", 0.3)
SPREAD = config.get("SPREAD", 1)
variable = config.get("variable", False)
NEIGHBORS = config.get("NEIGHBORS", 10)
pcs = config.get("pcs",30)


RES = config.get("resolution", 0.5)
n_top_genes = config.get("n_top_genes", 1000)
batch_method = config.get("batch_method", None)
n_comps = config.get("n_comps", 50)
recluster_resolution = config.get("recluster_resolution", 0.8)
recluster_n_top_genes = config.get("recluster_n_top_genes", 2000)
recluster_neighbors = config.get("recluster_neighbors", 15)
recluster_n_pcs = config.get("recluster_n_pcs", 30)
recluster_marker_method = config.get("recluster_marker_method", "wilcoxon")
recluster_min_pct = config.get("recluster_min_pct", 0.1)
recluster_logfc_threshold = config.get("recluster_logfc_threshold", 0.25)

if channel == "single_analysis" and batch_method!=None:
    L.info("your argument harmony just allowed setting in compare_analysis channel")
    batch_method = None

species = config.get("species", "human")
cluster_algorithm = config.get("cluster_algorithm", "leiden")
anno_algorithm_raw = str(config.get("anno_algorithm", "manual")).strip()
anno_algorithm_map = {
    "manual": "manual",
    "reannotation": "reannotation",
    "cell2location": "cell2Location",
    "rctd": "RCTD",
}
anno_algorithm_key = anno_algorithm_raw.lower()
if anno_algorithm_key not in anno_algorithm_map:
    raise ValueError(
        "anno_algorithm must be one of: manual, reannotation, cell2Location, RCTD; "
        f"got {anno_algorithm_raw!r}"
    )
anno_algorithm = anno_algorithm_map[anno_algorithm_key]
compare_algorithm_key = str(config.get("compare_algorithm", "DESeq2")).strip().lower()
if compare_algorithm_key in {"deseq2", "deseq"}:
    compare_algorithm = "DESeq2"
elif compare_algorithm_key == "edger":
    compare_algorithm = "edgeR"
else:
    raise ValueError("compare_algorithm must be DESeq2 or edgeR")
cell_focus = config.get("cell_focus", "all")
device = config.get("device", "cpu")
runpipe = config.get("runpipe", "pysenic")
counts_data = config.get("counts_data", "hgnc_symbol")
iterations = config.get("iterations", 1000)
threshold = config.get("threshold", 0.1)
try:
    workflow_threads = int(config.get("threads", 8))
except (TypeError, ValueError):
    raise ValueError("threads must be a positive integer")
if workflow_threads < 1:
    raise ValueError("threads must be >= 1")
pvalue = config.get("pvalue", 0.05)
sample_id = config.get("sample_type", "Normal")
microenvs_file_path = config.get("microenvs_file_path", "")
active_tf_path = config.get("active_tf_path", "")
degs_file_path = config.get("degs_file_path", "")
niche_col = config.get("niche_col", "spatial_cluster")
is_single_cell = config.get("is_single_cell", False)
cpdb_method = config.get("cpdb_method", "statistical")
n_clusters = config.get("n_clusters", 10)
output_name = config.get("output_name", "")
mt_threshold = config.get("mt_threshold", 50.0)
significance = config.get("significance", 0.05)
max_cluster = config.get("max_cluster", 10)
condition_col = config.get("condition_col", "condition")
sample_col = config.get("sample_col", "sample")
celltype_col = config.get("celltype_col", "celltype")
compare_celltype_col = config.get("compare_celltype_col", config.get("celltype_col", "celltype"))
compare_sample_col = config.get("compare_sample_col", config.get("sample_col", "sample"))
compare_condition_col = config.get("compare_condition_col", config.get("condition_col", "group"))
count_layer = config.get("count_layer", "counts")
min_replicates = config.get("min_replicates", 2)
min_cells_per_sample = config.get("min_cells_per_sample", 30)
min_total_counts_per_gene = config.get("min_total_counts_per_gene", 10)
de_top_n = config.get("de_top_n", 20)
_configured_compare_input_zarr = config.get("compare_input_zarr")
if _configured_compare_input_zarr is None or str(_configured_compare_input_zarr).strip().lower() in {
    "", "none", "null", "na", "nan"
}:
    compare_input_zarr = os.path.join(
        results_folder, "merge_data", "annotation", "concatenated_sdata.zarr"
    )
else:
    compare_input_zarr = str(_configured_compare_input_zarr)
compare_contrasts = config.get("compare_contrasts", [])
compare_gene_id_type = config.get("compare_gene_id_type", "auto")
compare_gene_symbol_col = config.get("compare_gene_symbol_col", "")
compare_go_ontology = str(config.get("compare_go_ontology", "BP")).upper()
compare_enrichment_top_n = int(config.get("compare_enrichment_top_n", 10))
sample_parameters = config.get("sample_parameters", {}) or {}
if not isinstance(sample_parameters, dict):
    raise ValueError("sample_parameters must be a YAML mapping keyed by sample_id")


def get_sample_paths(*args, **kwargs):
    kwargs.setdefault("sample_parameters", sample_parameters)
    kwargs.setdefault("default_bin_size", config.get("bin_size"))
    kwargs.setdefault("default_input_spec", config.get("input_spec", config.get("bin_size")))
    return _get_sample_paths(*args, **kwargs)


def get_stereoseq_input_spec_map(sample_list_file, channel):
    return _get_stereoseq_input_spec_map(
        sample_list_file,
        channel,
        sample_parameters=sample_parameters,
        default_input_spec=config.get("input_spec", config.get("bin_size")),
    )
liana_method = config.get("liana_method", "cellphonedb")
liana_resource_name = config.get("liana_resource_name", "consensus")
liana_expr_prop = config.get("liana_expr_prop", 0.1)
liana_min_cells = config.get("liana_min_cells", 5)
liana_use_raw = config.get("liana_use_raw", True)
liana_pvalue = config.get("liana_pvalue", 0.05)
liana_source_celltypes = config.get("liana_source_celltypes", "")
liana_target_celltypes = config.get("liana_target_celltypes", "")
liana_cell_pairs = config.get("liana_cell_pairs", "")
liana_pairs = config.get("liana_pairs", "")
cellcharter_col = config.get("cellcharter_col", "spatial_cluster")
cell_pairs = config.get("cell_pairs", "")
cell_type1 = config.get("cell_type1", "")
cell_type2 = config.get("cell_type2", "")
gene_family = config.get("gene_family", "")
cpdb_pathway = config.get("cpdb_pathway", "")
interaction_pairs = config.get("interaction_pairs", "")
cpdb_genes = config.get("cpdb_genes", "")
celltype = celltype_col
cellPhoneDB_input = config.get("cellPhoneDB_input", "")

geojson = "cell_segmentations.geojson"
image = "tissue_hires_image.png"
scale_factors = "scalefactors_json.json"
merscope_z_layers = config.get("merscope_z_layers", "None")
merscope_region_name = config.get("merscope_region_name", "None")
merscope_transcripts = config.get("merscope_transcripts", True)
merscope_cells_boundaries = config.get("merscope_cells_boundaries", True)
merscope_cells_table = config.get("merscope_cells_table", True)
merscope_mosaic_images = config.get("merscope_mosaic_images", True)

k_geom = config.get("k_geom", 15)
max_m = config.get("max_m", 1)
nbr_weight_decay = config.get("nbr_weight_decay", "scaled_gaussian")
lambda_list = config.get("lambda_list", [0.8])
banksy_n_comps = config.get("banksy_n_comps", 20)
banksy_resolution = config.get("banksy_resolution", [0.5])
banksy_num_nn = config.get("banksy_num_nn", 50)
banksy_max_features = config.get("banksy_max_features", 2000)
banksy_feature_col = config.get("banksy_feature_col", "highly_variable")
banksy_add_umap = config.get("banksy_add_umap", False)
banksy_plot_full = config.get("banksy_plot_full", False)
banksy_run_nonspatial = config.get("banksy_run_nonspatial", False)
banksy_plot_celltype_enrichment = config.get("banksy_plot_celltype_enrichment", True)
banksy_plot_max_points = config.get("banksy_plot_max_points", 200000)
banksy_sample_col = config.get("banksy_sample_col", "region")
banksy_selected_lambda = config.get("banksy_selected_lambda", "")
banksy_selected_resolution = config.get("banksy_selected_resolution", "")
banksy_seed = config.get("banksy_seed", 12345)

## downsample for spatialdata
sketch = config.get("sketch",False)
sample_rate = config.get("sample_rate",1.0)

INPUT_FILE=config.get("INPUT_FILE")
barcode=config.get('clusters')
max_x=config.get('max_x')
min_x=config.get("min_x")
max_y=config.get("max_y")
min_y=config.get("min_y")

feather_input = config.get("feather_input", "")

if run_type != "xenium":
    image_type = config.get("image_type", "hires")
    shape_type = config.get("shape_type", "cell_boundaries")
else:
    image_type = config.get("image_type", "morphology_focus")
    shape_type = config.get("shape_type", "cell_circle")
vis_mode = config.get("vis_mode", "auto")
consistent_option = ["integrate","preprocess","clustering","annotation_help","annotation"]


def build_compare_sample_key(sample_id, group_id):
    return f"{group_id}::{sample_id}"

def normalize_run_type_name(value):
    return re.sub(r"[\s_-]+", "", str(value).strip().lower())

def normalize_optional_cellchat_spec(value):
    if value is None:
        return ""
    text = str(value).strip()
    if text.lower() in {"", "none", "null", "na", "nan"}:
        return ""
    return text

def is_truthy_config(value):
    return str(value).strip().lower() in {"true", "1", "yes", "y"}

def prepare_cellchat_scale_specs(sample_ids, raw_specs, run_type):
    normalized_run_type = normalize_run_type_name(run_type)
    specs = [normalize_optional_cellchat_spec(spec) for spec in raw_specs]
    if len(specs) < len(sample_ids):
        specs.extend([""] * (len(sample_ids) - len(specs)))
    elif len(specs) > len(sample_ids):
        specs = specs[:len(sample_ids)]

    if normalized_run_type in {"visium", "visiumhd", "visiumsegment"}:
        missing_samples = [sample for sample, spec in zip(sample_ids, specs) if spec == ""]
        if missing_samples:
            sys.exit(
                "\nCellChat requires the third column of sample.txt to provide scalefactors_json "
                f"for run_type={run_type}. Missing samples: {', '.join(missing_samples)}"
            )
        return specs

    if normalized_run_type == "stereoseq":
        missing_samples = [sample for sample, spec in zip(sample_ids, specs) if spec == ""]
        if missing_samples:
            sys.exit(
                "\nCellChat requires the third column of sample.txt to provide Stereo-seq "
                f"bin_size or cellbin for run_type={run_type}. Missing samples: {', '.join(missing_samples)}"
            )
        for sample, spec in zip(sample_ids, specs):
            try:
                parse_stereoseq_input_spec(spec)
            except ValueError as exc:
                sys.exit(f"\nInvalid CellChat Stereo-seq spec for sample '{sample}': {exc}")
        return specs

    return []

def load_sample_input_dir_map(sample_list_file):
    sample_input_dir_map = {}
    if not os.path.isfile(sample_list_file):
        return sample_input_dir_map
    with open(sample_list_file) as sample_handle:
        next(sample_handle, None)
        for raw_line in sample_handle:
            raw_line = raw_line.strip()
            if not raw_line:
                continue
            parts = re.split(r'\s+', raw_line)
            if len(parts) < 2:
                continue
            sample_id = parts[0].strip()
            input_dir = parts[1].strip()
            sample_input_dir_map[sample_id] = input_dir
            if channel == "compare_analysis" and len(parts) >= 3:
                sample_input_dir_map[build_compare_sample_key(sample_id, parts[2].strip())] = input_dir
    return sample_input_dir_map

sample_input_dir_map = load_sample_input_dir_map(sample_list)
sample_stereoseq_input_spec_map = (
    get_stereoseq_input_spec_map(sample_list, channel)
    if run_type == "stereoseq" and option != "compare_stage"
    else {}
)

if  option in consistent_option and anno_algorithm == "manual":
    main_file = []
    samples = []
    bin_size = []
    group = []
    if channel == 'single_analysis':
        if run_type in runtype:
            samples, main_file = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
        elif run_type == "visium_HD":
            samples, main_file, bin_size = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
    else:
        if run_type == "stereoseq":
            samples, main_file, bin_size, group = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
        elif run_type in runtype:
            samples, main_file, group = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
        elif run_type == "visium_HD":
            samples, main_file, bin_size, group = get_sample_paths(sample_list,run_type,channel,option,require_non_empty=True, data_fold=data_fold)
    main_file = main_file[0]
elif option == "compare_stage" and runpipe != "cellchat":
    main_file = []
    samples = []
    bin_size = []
    group = []
    if channel == 'single_analysis':
        if run_type in runtype:
            samples, main_file = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
        elif run_type == "visium_HD":
            samples, main_file, bin_size = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
    else:
        if run_type == "stereoseq":
            samples, main_file, bin_size, group = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
        elif run_type in runtype:
            samples, main_file, group = get_sample_paths(sample_list, run_type,channel,option,require_non_empty=True, data_fold=data_fold)
        elif run_type == "visium_HD":
            samples, main_file, bin_size, group = get_sample_paths(sample_list,run_type,channel,option,require_non_empty=True, data_fold=data_fold)
    if len(main_file) > 0:
        main_file = main_file[0]
    unique_groups = pd.unique(group).tolist() if len(group) > 0 else []
    if len(unique_groups) > 0 and all(str(g).strip().lower().startswith("group") for g in unique_groups):
        L.info("compare_stage uses the third column of sample.txt directly for result naming. Please replace generic labels like Group1/Group2 with real biological group names.")
else:
    samples = []
    downstream_file = []
    reference=[]
    scale_factors_files = []
    cellchat_sample_names = []
    cellchat_scale_factors = []
    cellchat_output_name = None
    cellcharter_input = ""
    cellcharter_sample_id = sample_id
    banksy_input = ""
    banksy_sample_id = sample_id
    if option=="annotation":
        samples,downstream_file,reference = get_sample_paths(
            sample_list,
            run_type,
            channel,
            option,
            require_non_empty=False,
            data_fold=data_fold,
            annotation_reference_required=anno_algorithm != "reannotation",
        )
    else:
        samples,downstream_file,scale_factors_files= get_sample_paths(sample_list, run_type,channel,option,require_non_empty=False, data_fold=data_fold)
    if option == "compare_stage" and runpipe == "cellchat":
        if len(downstream_file) == 0:
            sys.exit("\ncompare_stage: 未提供 CellChat 比较所需的 rds 文件路径")
        if len(downstream_file) != 2:
            sys.exit(
                "\ncompare_stage CellChat requires exactly two condition-level .rds files. "
                "Use advance_analysis --runpipe=cellchat for one condition, or run explicit pairwise comparisons for more than two conditions."
            )
        bad = [p for p in downstream_file if (not os.path.isfile(p)) or (os.path.splitext(p)[1].lower() != ".rds")]
        if len(bad) > 0:
            L.info(f"以下路径不可用或非 .rds 文件: {bad}")
            sys.exit("\ncompare_stage: 请在 sample.txt 第二列提供有效的 CellChat .rds 路径")
    if option == "advance_analysis" and runpipe == "cellchat" and not is_truthy_config(config.get("cellchat_is_single_cell", config.get("is_single_cell", False))):
        scale_factors_files = prepare_cellchat_scale_specs(samples, scale_factors_files, run_type)
    if option == "advance_analysis" and runpipe == "cellchat" and channel == "compare_analysis":
        cellchat_sample_names = samples
        cellchat_scale_factors = scale_factors_files
        if len(downstream_file) > 0:
            derived_name = os.path.splitext(os.path.basename(downstream_file[0]))[0]
            cellchat_output_name = derived_name if derived_name else "concatenated_sdata"
        else:
            cellchat_output_name = "concatenated_sdata"
        samples = [cellchat_output_name]
    if option == "advance_analysis" and runpipe == "cellcharter":
        if len(downstream_file) == 0:
            sys.exit("\nadvance_analysis: 未提供 CellCharter 输入数据路径")
        cellcharter_input = downstream_file[0]
        if channel == "single_analysis" and len(samples) > 0:
            cellcharter_sample_id = samples[0]
        elif channel == "compare_analysis":
            derived_name = os.path.splitext(os.path.basename(downstream_file[0]))[0]
            cellcharter_sample_id = derived_name if derived_name else "concatenated_sdata"
    if option == "advance_analysis" and runpipe == "banksy":
        if len(downstream_file) == 0:
            sys.exit("\nadvance_analysis: 未提供 BANKSY 输入数据路径")
        banksy_input = downstream_file[0]
        if channel == "single_analysis" and len(samples) > 0:
            banksy_sample_id = samples[0]
        elif channel == "compare_analysis":
            derived_name = os.path.splitext(os.path.basename(downstream_file[0]))[0]
            banksy_sample_id = derived_name if derived_name else "concatenated_sdata"
    if option == "advance_analysis" and (runpipe == "cellPhoneDB" or runpipe == "pysenic" or runpipe == "liana"):
        if not cellPhoneDB_input:
            if len(downstream_file) == 0:
                sys.exit("\nadvance_analysis: 未提供 CellPhoneDB 输入数据路径")
            cellPhoneDB_input = downstream_file[0]
        if channel == "compare_analysis":
            cpdb_sample_id = "concatenated_sdata"
        elif len(samples) > 0:
            cpdb_sample_id = samples[0]
        else:
            cpdb_sample_id = sample_id
    print(samples,downstream_file,reference,scale_factors_files,"@@@@@@@@@@@")


subset_analysis_tasks = []
subset_task_lookup = {}
if option == "reclustering" or (option == "annotation" and anno_algorithm == "reannotation"):
    subset_analysis_tasks = build_subset_analysis_tasks(samples, downstream_file)
    subset_task_lookup = {
        (task["parent_sample"], task["subset_name"]): task
        for task in subset_analysis_tasks
    }


compare_task_lookup = {}
compare_output_paths = []
compare_result_root = os.path.join(results_folder, "merge_data", "compare_analysis", "compare_gene")


def _safe_compare_slug(value, used):
    original = str(value).strip()
    slug = re.sub(r"[^A-Za-z0-9_.-]+", "_", original).strip("_") or "unnamed"
    if slug in used and used[slug] != original:
        slug = f"{slug}_{hashlib.sha1(original.encode('utf-8')).hexdigest()[:8]}"
    used[slug] = original
    return slug


def _parse_compare_contrasts(raw_value, ordered_groups):
    parsed = []
    if raw_value in [None, "", []]:
        if len(ordered_groups) != 2:
            raise ValueError(
                "compare_contrasts is required unless sample.txt contains exactly two groups"
            )
        parsed = [{"comparison": ordered_groups[1], "reference": ordered_groups[0]}]
    elif isinstance(raw_value, str):
        for item in raw_value.split(","):
            item = item.strip()
            if not item:
                continue
            if ":" not in item:
                raise ValueError("compare_contrasts entries must use comparison:reference")
            comparison, reference_group = [part.strip() for part in item.split(":", 1)]
            parsed.append({"comparison": comparison, "reference": reference_group})
    elif isinstance(raw_value, (list, tuple)):
        for item in raw_value:
            if isinstance(item, dict):
                parsed.append(
                    {
                        "comparison": str(item.get("comparison", "")).strip(),
                        "reference": str(item.get("reference", "")).strip(),
                    }
                )
            elif isinstance(item, str) and ":" in item:
                comparison, reference_group = [part.strip() for part in item.split(":", 1)]
                parsed.append({"comparison": comparison, "reference": reference_group})
            else:
                raise ValueError("compare_contrasts must contain comparison/reference mappings")
    else:
        raise ValueError("compare_contrasts must be a YAML list or comma-separated string")

    available = set(ordered_groups)
    unique = []
    seen = set()
    for contrast in parsed:
        comparison = contrast["comparison"]
        reference_group = contrast["reference"]
        if not comparison or not reference_group or comparison == reference_group:
            raise ValueError("Each contrast requires different non-empty comparison and reference groups")
        missing = [group_name for group_name in (comparison, reference_group) if group_name not in available]
        if missing:
            raise ValueError("compare_contrasts contains groups absent from sample.txt: " + ", ".join(missing))
        key = (comparison, reference_group)
        if key not in seen:
            seen.add(key)
            unique.append(contrast)
    if not unique:
        raise ValueError("compare_contrasts did not define any comparisons")
    return unique


def _discover_compare_tasks():
    if option != "compare_stage" or runpipe == "cellchat" or channel != "compare_analysis":
        return
    if not os.path.isdir(compare_input_zarr):
        raise ValueError(f"compare_input_zarr is not an existing Zarr directory: {compare_input_zarr}")
    sample_rows = read_compare_sample_table(sample_list)
    ordered_groups = list(dict.fromkeys(row["group"] for row in sample_rows))
    contrasts = _parse_compare_contrasts(compare_contrasts, ordered_groups)

    import spatialdata as _compare_spatialdata

    sdata = _compare_spatialdata.read_zarr(compare_input_zarr)
    if "table" in sdata.tables:
        table = sdata.tables["table"]
    elif len(sdata.tables) > 0:
        table = sdata.tables[next(iter(sdata.tables.keys()))]
    else:
        raise ValueError("compare_input_zarr contains no SpatialData table")
    for required_column in (compare_sample_col, compare_celltype_col):
        if required_column not in table.obs.columns:
            raise ValueError(f"compare_input_zarr obs is missing required column: {required_column}")
    obs_sample = table.obs[compare_sample_col].astype(str)
    obs_celltype = table.obs[compare_celltype_col].astype(str)
    expected_samples = [row["sample_id"] for row in sample_rows]
    missing_samples = sorted(set(expected_samples) - set(obs_sample))
    if missing_samples:
        raise ValueError(
            f"sample.txt samples are absent from obs[{compare_sample_col!r}]: " + ", ".join(missing_samples)
        )

    available_targets = sorted(
        value for value in pd.unique(obs_celltype) if value not in {"", "nan", "None", "NA"}
    )
    requested_focus = str(cell_focus or "all").strip()
    if requested_focus.lower() in {"", "all", "none", "null"}:
        targets = available_targets
    else:
        targets = list(dict.fromkeys(item.strip() for item in requested_focus.split(",") if item.strip()))
        missing_targets = [target for target in targets if target not in set(available_targets)]
        if missing_targets:
            raise ValueError(
                "cell_focus values are absent from the selected cell-type column: "
                + ", ".join(missing_targets)
            )
    if not targets:
        raise ValueError("No cell types/regions are available for compare_analysis")

    target_slugs = {}
    contrast_slugs = {}
    sample_groups = {row["sample_id"]: row["group"] for row in sample_rows}
    used_targets = {}
    used_contrasts = {}
    for target in targets:
        target_slugs[target] = _safe_compare_slug(target, used_targets)
    for contrast in contrasts:
        label = f"{contrast['comparison']}_vs_{contrast['reference']}"
        contrast_slugs[label] = _safe_compare_slug(label, used_contrasts)

    selected_obs = obs_celltype.isin(targets)
    target_sample_counts = pd.crosstab(
        obs_celltype.loc[selected_obs],
        obs_sample.loc[selected_obs],
    )
    for target in targets:
        if target in target_sample_counts.index:
            target_counts = target_sample_counts.loc[target]
        else:
            target_counts = pd.Series(dtype="int64")
        for contrast in contrasts:
            comparison = contrast["comparison"]
            reference_group = contrast["reference"]
            valid_samples = [
                sample_id
                for sample_id in expected_samples
                if sample_groups[sample_id] in {comparison, reference_group}
                and int(target_counts.get(sample_id, 0)) >= int(min_cells_per_sample)
            ]
            group_counts = pd.Series(
                [sample_groups[sample_id] for sample_id in valid_samples]
            ).value_counts()
            if any(int(group_counts.get(group_name, 0)) < int(min_replicates) for group_name in (comparison, reference_group)):
                L.warning(
                    "Skipping compare task %s | %s vs %s: fewer than %s valid biological replicates in a group",
                    target,
                    comparison,
                    reference_group,
                    min_replicates,
                )
                continue
            target_slug = target_slugs[target]
            contrast_label = f"{comparison}_vs_{reference_group}"
            contrast_slug = contrast_slugs[contrast_label]
            compare_task_lookup[(target_slug, contrast_slug)] = {
                "celltype": target,
                "comparison": comparison,
                "reference": reference_group,
            }
            compare_output_paths.append(
                os.path.join(compare_result_root, compare_algorithm, target_slug, contrast_slug)
            )
    if not compare_output_paths:
        raise ValueError(
            f"No requested cell type/contrast has at least {min_replicates} "
            "valid biological replicates per group"
        )


_discover_compare_tasks()

def parameter_output(samples, option):
    outs = []
    if option == 'integrate':
        if channel == 'single_analysis':
            if run_type in runtype:
                outs += expand(os.path.join(results_folder, "{sample}",'integrate',"{sample}.zarr"), sample=samples)
            if run_type == "visium_HD":
                outs += expand(os.path.join(results_folder, "{sample}_{bin}um", 'integrate',"{sample}.zarr"), zip, sample=samples, bin=bin_size)
        if channel == "compare_analysis" and flag == 'merge':
            if run_type in runtype:
                outs += expand(os.path.join(results_folder, "{group}", "{sample}.zarr"), zip, sample=samples, group=group)
            if run_type == "visium_HD":
                outs += expand(os.path.join(results_folder, "{group}_{bin}um", "{sample}.zarr"), zip, sample=samples, bin=bin_size, group=group)
        if channel == "compare_analysis" and flag == 'ALL':
            outs.append(os.path.join(results_folder, "merge_data", "integrate", "concatenated_sdata.zarr"))
        return outs
    if option == "preprocess":
        if channel == 'single_analysis':
            if run_type in runtype:
                outs += expand(os.path.join(results_folder, "{sample}", 'preprocess', "filter_{sample}.zarr"), sample=samples)
            if run_type == "visium_HD":
                outs += expand(os.path.join(results_folder, "{sample}_{bin}um", 'preprocess', "filter_{sample}.zarr"), zip, sample=samples, bin=bin_size)
        if channel == "compare_analysis":
            outs.append(os.path.join(results_folder, "merge_data", 'preprocess', "filter_concatenated_sdata.zarr"))
        return outs
    if option == "annotation_help":
        if channel == 'single_analysis':
          if run_type == "visium_HD":
            return(expand(os.path.join(results_folder, "{sample}_{bin}um", 'clustering', 'kegg_data.csv'), zip, sample=samples, bin=bin_size))
          else:
            return(expand(os.path.join(results_folder, "{sample}", 'clustering', 'kegg_data.csv'), sample=samples))
        if channel == 'compare_analysis':
            return(os.path.join(results_folder, "merge_data", 'clustering', 'kegg_data.csv'))
    if option == "clustering" or (option == "annotation" and anno_algorithm == "manual"):
        if channel == 'single_analysis':
            if run_type in runtype:
                outs += expand(os.path.join(results_folder, "{sample}", option, "{sample}.zarr"), sample=samples)
            if run_type == "visium_HD":
                outs += expand(os.path.join(results_folder, "{sample}_{bin}um", option, "{sample}.zarr"), zip, sample=samples, bin=bin_size)
        if channel == "compare_analysis":
            outs.append(os.path.join(results_folder, "merge_data", option, "concatenated_sdata.zarr"))
        return outs
    if option == "reclustering":
        outs += [
            os.path.join(
                results_folder,
                task["parent_sample"],
                "reclustering",
                task["subset_name"],
                f"{task['subset_name']}.zarr",
            )
            for task in subset_analysis_tasks
        ]
        return outs
    if option == "annotation" and anno_algorithm == "reannotation":
        outs += [
            os.path.join(
                results_folder,
                task["parent_sample"],
                "reannotation",
                task["subset_name"],
                f"{task['subset_name']}.zarr",
            )
            for task in subset_analysis_tasks
        ]
        return outs
    if anno_algorithm != "manual" and option == "annotation":
        if anno_algorithm == "RCTD":
            outs += expand(os.path.join(results_folder, "{sample}", anno_algorithm, "{sample}.zarr"), sample=samples)
            return outs
        if anno_algorithm == "cell2Location":
            if channel == 'single_analysis':
                outs += expand(os.path.join(results_folder, "{sample}", "cell2Location", "{sample}.zarr"), sample=samples)
            if channel == "compare_analysis":
                outs.append(os.path.join(results_folder, "merge_data", "cell2Location", "concatenated_sdata.zarr"))
            return outs
        if channel == 'single_analysis':
            if run_type in runtype:
                outs += expand(os.path.join(results_folder, "{sample}", anno_algorithm, "{sample}.zarr"), sample=samples)
            if run_type == "visium_HD":
                outs += expand(os.path.join(results_folder, "{sample}", anno_algorithm, "{sample}.zarr"),sample=samples)
        if channel == "compare_analysis":
            outs.append(os.path.join(results_folder, "merge_data", anno_algorithm, "concatenated_sdata.zarr"))
        return outs
    if option == "compare_stage":
        if runpipe == "cellchat":
            return(cellchat_compare_output_dir)
        return compare_output_paths
    if option == "advance_analysis":
        if runpipe == "cellPhoneDB":
            if channel == 'single_analysis':
                outs += expand(os.path.join(results_folder, f"{cpdb_sample_id}", "cellPhoneDB_results", f"{cpdb_sample_id}_heatmap.png"))
            if channel == "compare_analysis":
                outs.append(os.path.join(results_folder, "merge_data", "cellPhoneDB_results", "concentrate_heatmap.png"))
        elif runpipe == "pysenic":
            outs.append(os.path.join(results_folder, "pysenic_results", f"{cpdb_sample_id}.aucell.loom"))
        elif runpipe == "liana":
            outs.append(os.path.join(results_folder,f"{cpdb_sample_id}","liana_output",f"{cpdb_sample_id}.zarr"))
        elif runpipe == "cellcharter":
            if channel == "compare_analysis":
                outs.append(os.path.join(results_folder, "merge_data", "cellcharter", f"{cellcharter_sample_id}_cellcharter.zarr"))
            else:
                outs.append(os.path.join(results_folder, f"{cellcharter_sample_id}", "cellcharter", f"{cellcharter_sample_id}_cellcharter.zarr"))
        elif runpipe == "banksy":
            if channel == "compare_analysis":
                outs.append(os.path.join(results_folder, "merge_data", "banksy", f"{banksy_sample_id}_banksy.zarr"))
            else:
                outs.append(os.path.join(results_folder, f"{banksy_sample_id}", "banksy", f"{banksy_sample_id}_banksy.zarr"))
        elif runpipe == "cellchat":
            cellchat_samples = samples
            if channel == "compare_analysis" and cellchat_output_name:
                cellchat_samples = [cellchat_output_name]
            outs += expand(os.path.join(results_folder, "{sample}", "cellchat", "{sample}_cellchat_network.png"), sample=cellchat_samples)
            outs += expand(os.path.join(results_folder, "{sample}", "cellchat", "{sample}_cellchat_network.pdf"), sample=cellchat_samples)
            outs += expand(os.path.join(results_folder, "{sample}", "cellchat", "{sample}_cellchat_stats.csv"), sample=cellchat_samples)
            outs += expand(os.path.join(results_folder, "{sample}", "cellchat", "{sample}_cellchat_lr.csv"), sample=cellchat_samples)
        return outs

def input_file(run_type):
    if run_type in runtype and run_type != "Merfish":
        return(os.path.join(data_fold, '{sample}', main_file))
    elif run_type == "Merfish":
        return(os.path.join(data_fold, "{sample}"))
    elif run_type == "visium_HD":
        return(os.path.join(data_fold, '{sample}', "binned_outputs", "square_{bin}um", main_file))
    elif run_type == "visium_segment":
        return os.path.join(data_fold, "{sample}", "segmented_outputs", main_file)

def get_output(run_type):
    if channel == 'single_analysis':
        if run_type == "visium_HD":
            return(directory(os.path.join(results_folder, "{sample}_{bin}um", "{sample}.zarr")))
        elif run_type in runtype:
            return(directory(os.path.join(results_folder, "{sample}", "{sample}.zarr")))
    elif channel == "compare_analysis":
        if run_type == "visium_HD":
            return(directory(os.path.join(results_folder, "{group}_{bin}um", "{sample}.zarr")))
        elif run_type in runtype:
            return(directory(os.path.join(results_folder, "{group}", "{sample}.zarr")))

def normal_file(run_type, segs):
    if channel == 'single_analysis':
        if run_type == "visium_HD":
            return((os.path.join(results_folder, "{sample}_{bin}um", segs, "{sample}.zarr")))
        elif run_type in runtype:
            return(os.path.join(results_folder, "{sample}", segs, "{sample}.zarr"))
    if channel == "compare_analysis":
        return((parameter_output(samples, segs)))

use_sample_filters = filter_list or (
    channel == "compare_analysis" and option in {"preprocess", "all"}
)

if not use_sample_filters:
    min_counts = config.get('min_counts', 200)
    min_cells = config.get('min_cells', 50)
    L.info(f"{min_counts},{min_cells}")
else:
    L.info("Per-sample filtering is enabled; resolving thresholds from YAML sample_parameters")
    default_min_counts = int(config.get("min_counts", 200))
    default_min_cells = int(config.get("min_cells", 50))
    default_mt_threshold = float(config.get("mt_threshold", 50.0))
    filter_dict = {}
    for current_sample in samples:
        overrides = sample_parameters.get(str(current_sample), {}) or {}
        sample_min_cells = int(overrides.get("min_cells", default_min_cells))
        sample_min_counts = int(overrides.get("min_counts", default_min_counts))
        sample_mt_threshold = float(overrides.get("mt_threshold", default_mt_threshold))
        if sample_min_cells < 1 or sample_min_counts < 1:
            raise ValueError(f"YAML filtering thresholds for {current_sample!r} must be >= 1")
        if not 0 <= sample_mt_threshold <= 100:
            raise ValueError(f"YAML mt_threshold for {current_sample!r} must be between 0 and 100")
        filter_dict[str(current_sample)] = [sample_min_cells, sample_min_counts, sample_mt_threshold]


flag = 'ALL'

rule all:
    input:
        parameter_output(samples, option)

if option == "integrate":
    include: "rules/integrate.smk"
    if channel == "compare_analysis":
        flag = 'merge'
        include: "rules/merge.smk"
elif option == "preprocess":
    include: "rules/preprocess.smk"
elif option == "clustering":
    include: "rules/cluster.smk"
elif option == "reclustering":
    include: "rules/reclustering.smk"
elif option == "annotation_help":
    include: "rules/annotation_help.smk"
elif option == "annotation":
    if anno_algorithm == "cell2Location":
        cell2location_input_spatial = downstream_file[0]
        cell2location_input_singlecell = reference[0]
        cell2location_spatial_by_sample = dict(zip(samples, downstream_file))
        cell2location_reference_by_sample = dict(zip(samples, reference))
        include: "rules/cell2Location_run.smk"
    elif anno_algorithm == "manual":
        anno_data = get_annotation(annotation_list, samples,channel,results_folder)
        include: "rules/manual.smk"
    elif anno_algorithm == "reannotation":
        anno_data = get_annotation(annotation_list, samples,channel,results_folder)
        include: "rules/reannotation.smk"
    elif anno_algorithm == "RCTD":
        input_spatial = downstream_file[0]
        input_singlecell = reference[0]
        rctd_spatial_by_sample = dict(zip(samples, downstream_file))
        rctd_reference_by_sample = dict(zip(samples, reference))
        include: "rules/RCTD.smk"
elif option == "compare_stage":
    if runpipe == "cellchat":
        include: "rules/compare_LR.smk"
    else:
        if channel == "compare_analysis":
            annotation_input = os.path.join(results_folder, "merge_data", "annotation", "concatenated_sdata.zarr")
        include: "rules/compare_gene.smk"
elif option == "advance_analysis":
    if runpipe == "cellPhoneDB":
        include: "rules/cellPhoneDB.smk"
    elif runpipe == "pysenic":
        input_pysenic = downstream_file[0]
        print(input_pysenic)
        include: "rules/py_senic.smk"
    elif runpipe == "liana":
        liana_inputs = downstream_file[0]
        include: "rules/run_liana.smk"
    elif runpipe == "cellcharter":
        include: "rules/run_cellcharter.smk"
    elif runpipe == "banksy":
        include: "rules/run_banksy.smk"
    elif runpipe == "cellchat":
        input_spatial = downstream_file[0]
        include: "rules/run_cellchat.smk"
elif option == "all":
    include: "rules/integrate.smk"
    if channel == "compare_analysis":
        flag = 'merge'
        include: "rules/merge.smk"
    include: "rules/preprocess.smk"
    include: "rules/cluster.smk"
    include: "rules/annotation_help.smk"
