## Cold-start benchmark for workflows 2-6, splitting load from compute.
## Same shape as split_bench.R: each method runs in a fresh R process which
## reports lib / load / compute separately, and the parent samples peak RSS.

D       <- "C:/Users/Matthew/Downloads/polars_and_duckdb_tutorial"
PQ      <- file.path(D, "fake_pcawg_all.parquet")
GENES   <- file.path(D, "gene_annotation.parquet")
KNOWN   <- file.path(D, "known_variants.parquet")
LONG    <- file.path(D, "rnaseq_counts_long.parquet")
WIDE    <- file.path(D, "rnaseq_counts_wide.parquet")
META    <- file.path(D, "sample_metadata.parquet")
KEYS    <- c("GENE101","GENE202","GENE303","GENE404","GENE505")
THREADS <- 8L
GB      <- 2^30
OUT     <- commandArgs(trailingOnly = TRUE)[1]

suppressPackageStartupMessages({ library(callr); library(ps) })

child <- function(method, paths, threads, keys) {
  Sys.setenv(POLARS_MAX_THREADS = as.character(threads))
  t0 <- proc.time()[["elapsed"]]

  pkgs <- switch(method,
    "wf2 DuckDB (range join)"    = c("DBI","duckdb"),
    "wf2 data.table (foverlaps)" = c("arrow","data.table"),
    "wf2 data.table (binned)"    = c("arrow","data.table"),
    "wf2 Polars (join + filter)" = c("polars"),
    "wf2 GenomicRanges"          = c("arrow","GenomicRanges"),
    "wf3 DuckDB"                 = c("DBI","duckdb"),
    "wf3 Polars (streaming)"     = c("polars"),
    "wf3 data.table"             = c("arrow","data.table"),
    "wf4 DuckDB"                 = c("DBI","duckdb"),
    "wf4 Polars (streaming)"     = c("polars"),
    "wf4 data.table (merge)"     = c("arrow","data.table"),
    "wf4 base R (merge)"         = c("arrow"),
    "wf4 data.table (packed key)" = c("arrow","data.table"),
    "wf5 base R matrix + BLAS"   = c("arrow"),
    "wf5 DuckDB (long) + BLAS"   = c("DBI","duckdb","data.table"),
    "wf6 DuckDB (one query)"     = c("DBI","duckdb"),
    "wf6 R in memory (naive)"    = c("arrow","dplyr")
  )
  for (p in pkgs) suppressPackageStartupMessages(
    library(p, character.only = TRUE, quietly = TRUE))
  if ("data.table" %in% pkgs) data.table::setDTthreads(threads)
  if ("arrow"      %in% pkgs) arrow::set_cpu_count(threads)
  t1 <- proc.time()[["elapsed"]]

  BINSZ <- 10000L
  st <- new.env()

  ## ---- LOAD ---------------------------------------------------------------
  switch(method,

    "wf2 DuckDB (range join)" = { st$cc <- DBI::dbConnect(duckdb::duckdb(),
        config = list(threads = as.character(threads))) },

    "wf2 data.table (foverlaps)" = {
      v <- data.table::as.data.table(as.data.frame(arrow::read_parquet(
             paths$pq, col_select = c("chromosome","position"))))
      v[, `:=`(start = as.integer(floor(position)), end = as.integer(floor(position)))]
      v[, position := NULL]
      g <- data.table::as.data.table(as.data.frame(arrow::read_parquet(paths$genes)))
      g <- g[, .(gene_id, chromosome, start = gene_start, end = gene_end)]
      data.table::setkey(g, chromosome, start, end)
      st$v <- v; st$g <- g },

    "wf2 data.table (binned)" = {
      g <- as.data.frame(arrow::read_parquet(paths$genes))
      nb <- (g$gene_end %/% BINSZ) - (g$gene_start %/% BINSZ) + 1L
      gt <- data.table::data.table(
        gene_id = rep(g$gene_id, nb), chromosome = rep(g$chromosome, nb),
        gene_start = rep(g$gene_start, nb), gene_end = rep(g$gene_end, nb),
        bin = unlist(mapply(function(a,b) a:b, g$gene_start %/% BINSZ,
                            g$gene_end %/% BINSZ, SIMPLIFY = FALSE), use.names = FALSE))
      data.table::setkey(gt, chromosome, bin)
      v <- data.table::as.data.table(as.data.frame(arrow::read_parquet(
             paths$pq, col_select = c("chromosome","position"))))
      v[, pos := as.integer(floor(position))][, position := NULL][, bin := pos %/% BINSZ]
      st$gt <- gt; st$v <- v },

    "wf2 Polars (join + filter)" = {
      st$pv <- polars::pl$scan_parquet(paths$pq)$select("chromosome","position")$
        with_columns(polars::pl$col("position")$floor()$cast(polars::pl$Int64)$alias("pos"))$
        drop("position")
      st$pg <- polars::pl$scan_parquet(paths$genes) },

    "wf2 GenomicRanges" = {
      v <- as.data.frame(arrow::read_parquet(
             paths$pq, col_select = c("chromosome","position")))
      st$vr <- GenomicRanges::GRanges(
        seqnames = as.character(v$chromosome),
        ranges   = IRanges::IRanges(start = as.integer(floor(v$position)), width = 1L))
      rm(v)
      g <- as.data.frame(arrow::read_parquet(paths$genes))
      st$gr <- GenomicRanges::GRanges(
        seqnames = as.character(g$chromosome),
        ranges   = IRanges::IRanges(start = g$gene_start, end = g$gene_end))
      st$gid <- g$gene_id },

    "wf3 DuckDB" = { st$cc <- DBI::dbConnect(duckdb::duckdb(),
        config = list(threads = as.character(threads))) },
    "wf3 Polars (streaming)" = { st$lf <- polars::pl$scan_parquet(paths$pq) },
    "wf3 data.table" = {
      st$dt <- data.table::as.data.table(as.data.frame(arrow::read_parquet(
        paths$pq, col_select = c("samplename","chromosome","position")))) },

    "wf4 DuckDB" = { st$cc <- DBI::dbConnect(duckdb::duckdb(),
        config = list(threads = as.character(threads))) },
    "wf4 Polars (streaming)" = {
      st$pv <- polars::pl$scan_parquet(paths$pq)$select("samplename","chromosome","position")$
        with_columns(polars::pl$col("position")$floor()$cast(polars::pl$Int64)$alias("pos"))$
        drop("position")
      st$pk <- polars::pl$scan_parquet(paths$known)$select("chromosome","pos","rsid") },
    "wf4 data.table (merge)" = {
      v <- data.table::as.data.table(as.data.frame(arrow::read_parquet(
             paths$pq, col_select = c("samplename","chromosome","position"))))
      v[, pos := as.integer(floor(position))][, position := NULL]
      k <- data.table::as.data.table(as.data.frame(arrow::read_parquet(
             paths$known, col_select = c("chromosome","pos","rsid"))))
      data.table::setkey(k, chromosome, pos)
      st$v <- v; st$k <- k },
    "wf4 base R (merge)" = {
      v <- as.data.frame(arrow::read_parquet(
             paths$pq, col_select = c("samplename","chromosome","position")))
      v$pos <- as.integer(floor(v$position)); v$position <- NULL
      k <- as.data.frame(arrow::read_parquet(
             paths$known, col_select = c("chromosome","pos","rsid")))
      st$v <- v; st$k <- k },

    "wf4 data.table (packed key)" = {
      v <- data.table::as.data.table(as.data.frame(arrow::read_parquet(
             paths$pq, col_select = c("samplename","chromosome","position"))))
      ## chromosome <= 22 and position < 2.5e8, so chr * 1e9 + pos is unique and
      ## sits well inside a double's 53 bits of exact integer range
      v[, k1 := as.numeric(chromosome) * 1e9 + floor(position)]
      v[, c("chromosome","position") := NULL]
      k <- data.table::as.data.table(as.data.frame(arrow::read_parquet(
             paths$known, col_select = c("chromosome","pos","rsid"))))
      k[, k1 := as.numeric(chromosome) * 1e9 + pos]
      k[, c("chromosome","pos") := NULL]
      data.table::setkey(k, k1)
      st$v <- v; st$k <- k },

    "wf5 base R matrix + BLAS" = {
      d <- as.data.frame(arrow::read_parquet(paths$wide))
      m <- as.matrix(d[, -1]); rownames(m) <- d$gene_id
      st$m <- m },
    "wf5 DuckDB (long) + BLAS" = { st$cc <- DBI::dbConnect(duckdb::duckdb(),
        config = list(threads = as.character(threads))) },

    "wf6 DuckDB (one query)" = { st$cc <- DBI::dbConnect(duckdb::duckdb(),
        config = list(threads = as.character(threads))) },
    "wf6 R in memory (naive)" = {
      st$sml_raw <- as.data.frame(arrow::read_parquet(
                      paths$pq, col_select = c("samplename","cluster_id")))
      st$meta   <- as.data.frame(arrow::read_parquet(paths$meta))
      st$counts <- as.data.frame(arrow::read_parquet(paths$counts))
      st$genes  <- as.data.frame(arrow::read_parquet(paths$genes)) }
  )
  invisible(length(ls(st)))
  t2 <- proc.time()[["elapsed"]]

  ## ---- COMPUTE ------------------------------------------------------------
  res <- switch(method,

    "wf2 DuckDB (range join)" = DBI::dbGetQuery(st$cc, sprintf("
        SELECT g.gene_id, COUNT(*) AS n_variants
        FROM (SELECT chromosome, CAST(FLOOR(position) AS BIGINT) AS pos
              FROM read_parquet('%s')) v
        JOIN read_parquet('%s') g
          ON v.chromosome = g.chromosome AND v.pos BETWEEN g.gene_start AND g.gene_end
        GROUP BY ALL", paths$pq, paths$genes)),

    "wf2 data.table (foverlaps)" = {
      ov <- data.table::foverlaps(st$v, st$g, type = "within", nomatch = NULL)
      ov[, .(n_variants = .N), by = gene_id] },

    "wf2 data.table (binned)" = {
      m <- st$gt[st$v, on = .(chromosome, bin), allow.cartesian = TRUE, nomatch = NULL]
      m[pos >= gene_start & pos <= gene_end, .(n_variants = .N), by = gene_id] },

    "wf2 Polars (join + filter)" = as.data.frame(
      st$pv$join(st$pg, on = "chromosome", how = "inner")$
        filter((polars::pl$col("pos") >= polars::pl$col("gene_start")) &
               (polars::pl$col("pos") <= polars::pl$col("gene_end")))$
        group_by("gene_id")$agg(polars::pl$len()$alias("n_variants"))$
        collect(engine = "streaming")),

    "wf2 GenomicRanges" = {
      hits <- GenomicRanges::findOverlaps(st$vr, st$gr, type = "within")
      tab  <- table(S4Vectors::subjectHits(hits))
      data.frame(gene_id    = st$gid[as.integer(names(tab))],
                 n_variants = as.integer(tab), stringsAsFactors = FALSE) },

    "wf3 DuckDB" = DBI::dbGetQuery(st$cc, sprintf("
        SELECT samplename, chromosome,
               CAST(FLOOR(position / 1000000) AS INTEGER) AS bin, COUNT(*) AS n
        FROM read_parquet('%s') GROUP BY ALL", paths$pq)),
    "wf3 Polars (streaming)" = as.data.frame(st$lf$
        with_columns((polars::pl$col("position")/1e6)$floor()$cast(polars::pl$Int32)$alias("bin"))$
        group_by("samplename","chromosome","bin")$agg(polars::pl$len()$alias("n"))$
        collect(engine = "streaming")),
    "wf3 data.table" = {
      dt <- st$dt
      dt[, bin := as.integer(floor(position/1000000))]
      dt[, .(n = .N), by = .(samplename, chromosome, bin)] },

    "wf4 DuckDB" = DBI::dbGetQuery(st$cc, sprintf("
        SELECT v.samplename, COUNT(*) AS n_variants, COUNT(k.rsid) AS n_known,
               COUNT(*) - COUNT(k.rsid) AS n_novel
        FROM (SELECT samplename, chromosome, CAST(FLOOR(position) AS BIGINT) AS pos
              FROM read_parquet('%s')) v
        LEFT JOIN read_parquet('%s') k ON k.chromosome = v.chromosome AND k.pos = v.pos
        GROUP BY ALL", paths$pq, paths$known)),
    "wf4 Polars (streaming)" = as.data.frame(
      st$pv$join(st$pk, on = c("chromosome","pos"), how = "left")$
        group_by("samplename")$
        agg(polars::pl$len()$alias("n_variants"),
            polars::pl$col("rsid")$count()$alias("n_known"))$
        with_columns((polars::pl$col("n_variants")-polars::pl$col("n_known"))$alias("n_novel"))$
        collect(engine = "streaming")),
    "wf4 data.table (merge)" = {
      m <- st$k[st$v, on = .(chromosome, pos)]
      m[, .(n_variants = .N, n_known = sum(!is.na(rsid)),
            n_novel = sum(is.na(rsid))), by = samplename] },
    "wf4 data.table (packed key)" = {
      m <- st$k[st$v, on = .(k1)]
      m[, .(n_variants = .N, n_known = sum(!is.na(rsid)),
            n_novel = sum(is.na(rsid))), by = samplename] },

    "wf4 base R (merge)" = {
      m <- merge(st$v, st$k, by = c("chromosome","pos"), all.x = TRUE)
      n_tot <- aggregate(list(n_variants = m$pos),
                         by = list(samplename = m$samplename), FUN = length)
      n_kn  <- aggregate(list(n_known = !is.na(m$rsid)),
                         by = list(samplename = m$samplename), FUN = sum)
      r <- merge(n_tot, n_kn, by = "samplename")
      r$n_novel <- r$n_variants - r$n_known
      r },

    "wf5 base R matrix + BLAS" = {
      m <- st$m
      lib  <- colSums(m)
      keep <- rowSums(m > 0) > ncol(m) * 0.5
      cpm  <- log2(t(t(m[keep, ]) / lib) * 1e6 + 1)
      v    <- apply(cpm, 1, var)
      top  <- order(v, decreasing = TRUE)[seq_len(min(2000L, sum(keep)))]
      cm   <- cor(cpm[top, ])
      data.frame(samplename = colnames(cm), mean_cor = colMeans(cm),
                 lib_size = lib[colnames(cm)]) },
    "wf5 DuckDB (long) + BLAS" = {
      long <- paths$long
      cpm <- DBI::dbGetQuery(st$cc, sprintf("
        WITH lib AS (SELECT samplename, SUM(count) AS total FROM read_parquet('%s') GROUP BY ALL),
             keep AS (SELECT gene_id FROM read_parquet('%s') GROUP BY ALL
                      HAVING SUM(CASE WHEN count > 0 THEN 1 ELSE 0 END) > 0.5 * COUNT(*)),
             norm AS (SELECT c.gene_id, c.samplename,
                             log2(1e6 * c.count / l.total + 1) AS lcpm
                      FROM read_parquet('%s') c JOIN lib l USING (samplename)
                      SEMI JOIN keep k ON k.gene_id = c.gene_id),
             topv AS (SELECT gene_id FROM norm GROUP BY ALL
                      ORDER BY var_samp(lcpm) DESC LIMIT 2000)
        SELECT n.gene_id, n.samplename, n.lcpm
        FROM norm n SEMI JOIN topv t ON t.gene_id = n.gene_id", long, long, long))
      m <- data.table::dcast(data.table::as.data.table(cpm),
                             gene_id ~ samplename, value.var = "lcpm")
      mm <- as.matrix(m[, -1]); cm <- cor(mm)
      lib <- DBI::dbGetQuery(st$cc, sprintf(
        "SELECT samplename, SUM(count) AS total FROM read_parquet('%s') GROUP BY ALL", long))
      data.frame(samplename = colnames(cm), mean_cor = colMeans(cm),
                 lib_size = lib$total[match(colnames(cm), lib$samplename)]) },

    "wf6 DuckDB (one query)" = DBI::dbGetQuery(st$cc, sprintf("
        WITH sml AS (SELECT samplename,
               (COUNT(*) - COUNT(*) FILTER (cluster_id = mn))::DOUBLE / COUNT(*) AS sMF
             FROM (SELECT samplename, cluster_id,
                          MIN(cluster_id) OVER (PARTITION BY samplename) mn
                   FROM read_parquet('%s')) GROUP BY ALL),
        grp AS (SELECT s.samplename, m.cancer_type,
               CASE WHEN s.sMF > MEDIAN(s.sMF) OVER (PARTITION BY m.cancer_type)
                    THEN 'high' ELSE 'low' END AS sml_group
             FROM sml s JOIN read_parquet('%s') m USING (samplename)),
        lib AS (SELECT samplename, SUM(count) AS total FROM read_parquet('%s') GROUP BY ALL),
        keys AS (SELECT gene_id, symbol FROM read_parquet('%s') WHERE symbol IN ('%s'))
        SELECT k.symbol, c.samplename, g.cancer_type, g.sml_group,
               1e6 * c.count / l.total AS cpm
        FROM read_parquet('%s') c JOIN keys k USING (gene_id)
        JOIN grp g USING (samplename) JOIN lib l USING (samplename)",
        paths$pq, paths$meta, paths$counts, paths$genes,
        paste(keys, collapse = "','"), paths$counts)),

    "wf6 R in memory (naive)" = {
      sml <- dplyr::mutate(dplyr::summarize(st$sml_raw,
               mn = min(cluster_id), n = dplyr::n(),
               k = sum(cluster_id == min(cluster_id)), .by = samplename),
               sMF = (n - k) / n)
      grp <- dplyr::mutate(dplyr::inner_join(sml, st$meta, by = "samplename"),
               sml_group = ifelse(sMF > stats::median(sMF), "high", "low"), .by = cancer_type)
      lib <- dplyr::summarize(st$counts, total = sum(count), .by = samplename)
      kk  <- st$genes[st$genes$symbol %in% keys, c("gene_id","symbol")]
      dplyr::transmute(
        dplyr::inner_join(dplyr::inner_join(dplyr::inner_join(st$counts, kk, by = "gene_id"),
          grp[, c("samplename","cancer_type","sml_group")], by = "samplename"),
          lib, by = "samplename"),
        symbol, samplename, cancer_type, sml_group, cpm = 1e6 * count / total) }
  )

  res <- as.data.frame(res)
  invisible(nrow(res))
  t3 <- proc.time()[["elapsed"]]
  if (!is.null(st$cc)) try(DBI::dbDisconnect(st$cc, shutdown = TRUE), silent = TRUE)

  list(nrow = nrow(res), digest = {
         num <- vapply(res, function(cl) if (is.numeric(cl)) sum(as.numeric(cl)) else NA_real_,
                       numeric(1))
         round(sum(num, na.rm = TRUE), 3) },
       lib_s = t1 - t0, load_s = t2 - t1, compute_s = t3 - t2)
}

PATHS <- list(pq = PQ, genes = GENES, known = KNOWN, long = LONG, wide = WIDE,
              meta = META, counts = LONG)

run_one <- function(method) {
  t_launch <- proc.time()[["elapsed"]]
  p <- callr::r_bg(child, args = list(method = method, paths = PATHS,
                                      threads = THREADS, keys = KEYS), supervise = TRUE)
  peak <- 0; n <- 0L; h <- NULL
  repeat {
    if (is.null(h)) h <- tryCatch(ps::ps_handle(p$get_pid()), error = function(e) NULL)
    if (!is.null(h)) {
      m <- tryCatch(ps::ps_memory_info(h)[["rss"]], error = function(e) NA_real_)
      if (!is.na(m)) { n <- n + 1L; if (m > peak) peak <- m }
    }
    if (!p$is_alive()) break
    Sys.sleep(0.02)
  }
  total <- proc.time()[["elapsed"]] - t_launch
  out <- tryCatch(p$get_result(), error = function(e) NULL)
  if (is.null(out)) {
    msg <- paste(utils::tail(p$read_all_error_lines(), 4), collapse = " | ")
    message("  FAILED: ", substr(msg, 1, 220))
    return(data.frame(method = method, total_s = NA, startup_s = NA, lib_s = NA,
      load_s = NA, compute_s = NA, peak_gb = NA, nrow = NA, digest = NA,
      stringsAsFactors = FALSE))
  }
  startup <- max(0, total - out$lib_s - out$load_s - out$compute_s)
  data.frame(method = method, total_s = total, startup_s = startup, lib_s = out$lib_s,
    load_s = out$load_s, compute_s = out$compute_s, peak_gb = peak / GB,
    nrow = out$nrow, digest = out$digest, stringsAsFactors = FALSE)
}

METHODS <- c(
  "wf3 DuckDB", "wf3 Polars (streaming)", "wf3 data.table",
  "wf6 DuckDB (one query)", "wf6 R in memory (naive)",
  "wf5 base R matrix + BLAS", "wf5 DuckDB (long) + BLAS",
  "wf2 DuckDB (range join)", "wf2 data.table (foverlaps)", "wf2 data.table (binned)",
  "wf2 GenomicRanges",
  "wf4 DuckDB", "wf4 Polars (streaming)", "wf4 data.table (merge)",
  "wf4 data.table (packed key)",
  "wf2 Polars (join + filter)", "wf4 base R (merge)")

invisible(readBin(PQ, "raw", n = file.size(PQ))); gc(full = TRUE)

rows <- list()
for (m in METHODS) {
  message("== ", m)
  r <- run_one(m)
  message(sprintf("   total %.2fs = startup %.2f + libs %.2f + load %.2f + compute %.2f | peak %.2f GB | rows %s digest %s",
                  r$total_s, r$startup_s, r$lib_s, r$load_s, r$compute_s, r$peak_gb,
                  format(r$nrow, big.mark = ","), format(r$digest)))
  rows[[length(rows)+1L]] <- r
  utils::write.csv(do.call(rbind, rows), OUT, row.names = FALSE)
  gc(full = TRUE)
}
print(do.call(rbind, rows), row.names = FALSE, digits = 3)
