suppressPackageStartupMessages(library(argparse))

parser <- ArgumentParser(description = "Convert Seurat object to AnnData for SCENIC+")
parser$add_argument("-i", "--input_rds", type = "character", required = TRUE,
                    help = "Path to the input Seurat RDS file (containing both RNA and ATAC assays)")
parser$add_argument("-m", "--meta_path", type = "character", default = NA,
                    help = "Path to external metadata file (TSV, optional)")
parser$add_argument("-c", "--celltype_col", type = "character",
                    default = "CellAnnotation",
                    help = "Column name for cell type annotation in metadata")
parser$add_argument("-o", "--output_dir", type = "character",
                    default = "./output",
                    help = "Output directory for the converted AnnData files")
parser$add_argument("-d", "--downsample", type = "logical",
                    default = FALSE,
                    help = "Whether to downsample cells per cell type")
parser$add_argument("-n", "--downsample_num", type = "integer",
                    default = 500,
                    help = "Number of cells to downsample per cell type")

args <- parser$parse_args()

library(reticulate)
library(Seurat)
library(anndata)
library(Matrix)

SEURAT_RDS_PATH  <- args$input_rds
META_PATH        <- args$meta_path
OUTPUT_DIR       <- args$output_dir
GROUP_COL        <- "orig.ident"
SAMPLE_COL       <- "Sample"
RNA_ASSAY_NAME   <- "RNA"
ATAC_ASSAY_NAME  <- "ATAC"
RNA_OUTPUT_FILE  <- "scRNA.h5ad"
ATAC_OUTPUT_FILE <- "scATAC.h5ad"

if (!dir.exists(OUTPUT_DIR)) {
  dir.create(OUTPUT_DIR, recursive = TRUE)
}

cat(sprintf("[Step 0] Reading Seurat object from: %s\n", SEURAT_RDS_PATH))
seurat_obj <- readRDS(SEURAT_RDS_PATH)

if (!is.na(META_PATH) && file.exists(META_PATH)) {
  cat(sprintf("[Step 0] Reading external metadata from: %s\n", META_PATH))
  external_meta <- read.table(META_PATH, header = TRUE, sep = "\t",
                              row.names = 1, check.names = FALSE)
  external_meta <- external_meta[, c(args$celltype_col, GROUP_COL, SAMPLE_COL)]
  cat("[Step 0] Merging external metadata into Seurat object...\n")
  seurat_obj <- AddMetaData(seurat_obj, metadata = external_meta)
} else {
  cat("[Step 0] No external metadata provided. Using existing metadata.\n")
}

if (isTRUE(args$downsample)) {
  cat(sprintf("[Step 0] Downsampling to %d cells per cell type...\n", args$downsample_num))
  set.seed(42)
  cells_to_keep <- unlist(lapply(
    split(colnames(seurat_obj), seurat_obj@meta.data[[args$celltype_col]]),
    function(cells) {
      if (length(cells) > args$downsample_num) sample(cells, args$downsample_num) else cells
    }
  ))
  seurat_obj <- subset(seurat_obj, cells = cells_to_keep)
  cat(sprintf("[Step 0] After downsampling: %d cells remaining.\n", ncol(seurat_obj)))
}

cat(sprintf("[Step 0] Seurat object: %d cells, %d features, %d assays\n",
            ncol(seurat_obj), nrow(seurat_obj), length(seurat_obj@assays)))
print(seurat_obj)

convert_assay_to_anndata <- function(obj, assay_name, file_name) {
  cat(sprintf("[Step 0] Converting assay: %s\n", assay_name))

  if (!assay_name %in% names(obj@assays)) {
    stop(sprintf("Assay '%s' not found in Seurat object!", assay_name))
  }

  counts_matrix <- tryCatch({
    GetAssayData(obj, assay = assay_name, layer = "counts")
  }, error = function(e) {
    cat("[Step 0] Layer 'counts' not found, trying slot 'counts'...\n")
    GetAssayData(obj, assay = assay_name, slot = "counts")
  })

  counts_matrix <- t(counts_matrix)

  if (assay_name == ATAC_ASSAY_NAME) {
    cat("[Step 0] Formatting ATAC peak names to SCENIC+ standard (chr:start-end)...\n")
    feature_names <- colnames(counts_matrix)
    feature_names <- gsub("^([^:-]+)-([0-9]+)-([0-9]+)$", "\\1:\\2-\\3", feature_names)
    feature_names <- gsub("^([^:-]+)_([0-9]+)_([0-9]+)$", "\\1:\\2-\\3", feature_names)
    colnames(counts_matrix) <- feature_names
  }

  metadata <- obj@meta.data

  adata <- AnnData(
    X = counts_matrix,
    obs = metadata,
    var = data.frame(row.names = colnames(counts_matrix))
  )

  adata$raw <- adata

  output_path <- file.path(OUTPUT_DIR, file_name)
  cat(sprintf("[Step 0] Saving AnnData to: %s\n", output_path))
  write_h5ad(adata, output_path, compression = "gzip")
  cat(sprintf("[Step 0] Done. Shape: %d cells × %d features\n", nrow(counts_matrix), ncol(counts_matrix)))
}

convert_assay_to_anndata(seurat_obj, RNA_ASSAY_NAME,  RNA_OUTPUT_FILE)
convert_assay_to_anndata(seurat_obj, ATAC_ASSAY_NAME, ATAC_OUTPUT_FILE)

cat("\n[Step 0] Conversion completed successfully!\n")
cat(sprintf("[Step 0] Output files:\n"))
cat(sprintf("[Step 0]   %s\n", file.path(OUTPUT_DIR, RNA_OUTPUT_FILE)))
cat(sprintf("[Step 0]   %s\n", file.path(OUTPUT_DIR, ATAC_OUTPUT_FILE)))
