#!/usr/bin/env Rscript
# ============================================================
# Compare Bulk (DESeq2/RUVr) vs Single-Cell (Seurat/Wilcox) DEGs
# Bulk:  Old vs Young, k=3, sig only
# SC:    CM Old vs Young, sig only (p_val_adj<0.01, |log2FC|>1)
# ============================================================

library(dplyr)

# ============================================================
# 0. Gene ID mapping  (Ensembl -> Symbol from count table)
# Both DEG files and count tables originate from the same GTEx
# count matrix, so versioned IDs (e.g. ENSG00000237973.1) are
# identical — direct join on ensembl_full, no version stripping.
# Use sex-specific count tables for female / male respectively.
# ============================================================
load_map <- function(map_file) {
  m <- read.table(map_file, header = TRUE, sep = "\t",
                  check.names = FALSE)[, 1:2]
  colnames(m) <- c("ensembl_full", "symbol")
  m[!duplicated(m$ensembl_full), ]
}

map_male   <- load_map("/BRC/yan/heart/getx_LV/male_Heart_-_Left_Ventricle/Heart_-_Left_Ventricle_sex1_filter.txt")
map_female <- load_map("/BRC/yan/heart/getx_LV/female_Heart_-_Left_Ventricle/Heart_-_Left_Ventricle_sex2_filter.txt")

# ============================================================
# helper: load bulk DEG, map to symbol, keep sig
# ============================================================
load_bulk <- function(path, id_map) {
  df <- read.table(path, header = TRUE, sep = "\t", check.names = FALSE)
  # direct join on versioned Ensembl ID (same source as count table)
  df <- left_join(df, id_map, by = c("gene" = "ensembl_full"))
  df$direction_bulk <- df$sig   # "UP" or "DOWN"
  df[!is.na(df$symbol), c("symbol", "direction_bulk",
                           "log2FoldChange", "padj")]
}

# helper: load SC DEG, assign direction
load_sc <- function(path) {
  df <- read.csv(path, check.names = FALSE)
  df$direction_sc <- ifelse(df$avg_log2FC > 0, "UP", "DOWN")
  df[, c("gene", "direction_sc", "avg_log2FC", "p_val_adj")]
}

# ============================================================
# helper: compare and report
# ============================================================
compare_degs <- function(bulk_path, sc_path, label, id_map) {
  bulk <- load_bulk(bulk_path, id_map)
  sc   <- load_sc(sc_path)

  cat("\n", strrep("=", 60), "\n", sep = "")
  cat(" ", label, "\n")
  cat(strrep("=", 60), "\n")
  cat("  Bulk DEGs  :", nrow(bulk),
      " (UP:", sum(bulk$direction_bulk == "UP"),
      " DOWN:", sum(bulk$direction_bulk == "DOWN"), ")\n")
  cat("  SC DEGs    :", nrow(sc),
      " (UP:", sum(sc$direction_sc == "UP"),
      " DOWN:", sum(sc$direction_sc == "DOWN"), ")\n")

  shared <- inner_join(bulk, sc, by = c("symbol" = "gene"))

  cat("\n  Shared genes:", nrow(shared), "\n\n")

  if (nrow(shared) > 0) {
    shared$direction_match <- shared$direction_bulk == shared$direction_sc

    # direction summary
    dir_tbl <- shared %>%
      count(direction_bulk, direction_sc) %>%
      mutate(concordant = direction_bulk == direction_sc)

    cat("  Direction breakdown:\n")
    print(dir_tbl)

    n_agree <- sum(shared$direction_match)
    n_dis   <- sum(!shared$direction_match)
    cat("\n  Concordant (same direction) :", n_agree, "\n")
    cat("  Discordant (opposite direction):", n_dis, "\n")

    cat("\n  --- Shared gene details ---\n")
    shared_out <- shared %>%
      arrange(direction_bulk, direction_sc, symbol) %>%
      select(symbol, direction_bulk, direction_sc,
             log2FC_bulk = log2FoldChange,
             log2FC_sc   = avg_log2FC,
             padj_bulk   = padj,
             padj_sc     = p_val_adj)
    print(as.data.frame(shared_out))

    # save to file
    out_path <- file.path("/BRC/yan/heart/analysis",
                          paste0("shared_DEG_", gsub("[^A-Za-z0-9]", "_", label), ".csv"))
    write.csv(shared_out, out_path, row.names = FALSE)
    cat("\n  Saved to:", out_path, "\n")
  }

  invisible(shared)
}

# ============================================================
# FEMALE
# ============================================================
res_female <- compare_degs(
  bulk_path = "/BRC/yan/heart/analysis/DEG_LV/female_v2/sex2_v2_DEG_old_vs_young_RUVr_k3_filter.txt",
  sc_path   = "/BRC/yan/heart/amina_res/DEG/CM_female_Old_vs_Young_sig.csv",
  label     = "FEMALE  |  Bulk vs SC  |  CM  Old vs Young",
  id_map    = map_female
)

# ============================================================
# MALE
# ============================================================
res_male <- compare_degs(
  bulk_path = "/BRC/yan/heart/analysis/DEG_LV/male_v2/sex1_v2_DEG_old_vs_young_RUVr_k3_filter.txt",
  sc_path   = "/BRC/yan/heart/amina_res/DEG/CM_male_Old_vs_Young_sig.csv",
  label     = "MALE    |  Bulk vs SC  |  CM  Old vs Young",
  id_map    = map_male
)

cat("\n[Done]\n")
