# TODO: Add comment
# 
# Author: E.Korsching  2024 versions
###############################################################################



# install.packages("RRNA")


# customize data
f2cust <- function(x, norm=T, ctl2sample=NULL, ctl.prim.sage=NULL, ctl.val.sage=NULL, ctl.val.gtl=NULL ){
	# make integer and normalize or not etc.
	# primary samples: sage (1), validation samples: sage (2) + gtl (3) vector
	# ctl.prim.sage, ctl.val.sage, ctl.val.gtl : positions
	# ctl2sample: vec 1,2,3 or 0 for controls+rest
	cat("\n\n column names : \n", dimnames(x)[[2]] )
	
	if(norm){
		# normalization according to used media:  Sage or GTL  see Excel table
		sage.prime <- apply(x[,ctl.prim.sage], 1, median)		# we take median  of controls (sage)
		sage.val <- apply(x[,ctl.val.sage], 1, median)		# we take median  of controls (sage)
		gtl.val <- apply(x[,ctl.val.gtl], 1, median)		# we take median  of controls (gtl)
		
		# subtract for each column the right base line
		cat("\n\n ctl2sample \n", ctl2sample)
		for(i in 1:length(ctl2sample)){
			if(ctl2sample[i]==1){
				x[,i] <- x[,i]-sage.prime		# we substract
			}
			if(ctl2sample[i]==2){
				x[,i] <- x[,i]-sage.val
			}
			if(ctl2sample[i]==3){
				x[,i] <- x[,i]-gtl.val
			}
		}
	}
	
	# make integer
	f1 <- function(x){ as.integer(floor(x)) }		# make integer: conservative: floor
	a <- row.names(x)
	x <- apply(x, 2, f1)			# matrix returned
	row.names(x) <- a				# restore row names
	# remove <0
	total <- dim(x)[[1]]*dim(x)[[2]]
	cat("\n\n number x<0 : ? : ", sum(x<0) , round((sum(x<0)/total)*100, 4) , "%")
	x[x<0] <- 0
	cat("\n number x [0..1] : ? : ", sum(x>0&x<1) , round((sum(x>0&x<1)/total)*100, 4) , "%")
	cat("\n number x [0]    : ? : ", sum(x==0) , round((sum(x==0)/total)*100, 4) , "%")
	# remove 0 lines
	cat("\n rows before     : ? : ", nrow(x))
	x <- x[rowSums(x)>0,]				# keep if rowSum > 0
	cat("\n rows after      : ? : ", nrow(x), " sum(row)==0 removed")
	# remove 0 values
	x[x<1 & x>=0] <- as.integer(1)
	return(x)
}


de.wrap <- function(r.vec, d.vec, sn.list, rap, ran, wpath, test="nbW", stati="de_overview"){
	# DE wrapper function
	# r.vec: result name vector
	# d.vec: data name vector
	# sn.list: selected column names, list
	# rap: column range numbers of affected samples, list of vectors
	# ran: column range numbers of normal samples, list of vectors
	# wpath: relative output path,   test: one of  DESeq > DESeqMean > nbW > LRT
	rlen <- length(r.vec)
	dlen <- length(d.vec)
	snlen <- length(sn.list)
	raplen <- length(rap)
	ranlen <- length(ran)
	dir.create(file.path(wpath, test))
	wpath <- paste(wpath,test,"/",sep="")
	#paste(format(Sys.time(), "0_%Y_%m_%d_%M%S_"),"de_count.txt",sep="")
	write.table(data.frame("START"), file=paste(wpath,stati,".tsv",sep=""), append=F, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	
	k <- 1	# results
	for(i in 1:dlen){	# data elements
		blo.sp <- unlist(strsplit(d.vec[i],"$",fixed=T))		# [1,2,3] [1]
		blo.sp.l <- length(blo.sp)
		blo <- get(blo.sp[1], envir=.GlobalEnv)
		for(j in 1:snlen){	# calculation(s)
			#blsn <- unlist(mget(sn.list[[j]], envir=.GlobalEnv))		# changed, is getting names directly now
			blsn <- sn.list[[j]]
			blrap <- rap[[j]]
			blran <- ran[[j]]
			blout <- r.vec[k]
			
			if(blo.sp.l==3){
				blo.n <- blo[[blo.sp[2]]][[blo.sp[3]]]
				dat <- de.pre(blo.n[,blsn], wpath=wpath, sfile=stati)	# undefined columns selected
			}else if(blo.sp.l==1){
				dat <- de.pre(blo[,blsn], wpath=wpath, sfile=stati)
			}
			len <- length(c(blrap,blran))
			cat("\nlen",len)
			a1 <- vector("character",len)
			a1[blrap] <- "A"
			a1[blran] <- "N"
			
			out <- de.fn(dat,
					a1=a1,
					wpath=wpath,
					wfile=blout,
					sfile=stati,
					anno=NULL, sel=c("condition","A","N"), sign=F, test=test, ret=T)
			assign(blout, out, pos=1)
			k <- k+1
		}
	}
	return()
}


de.pre <- function(x, wpath, sfile){
	# customize and remove empty rows
	# x: data frame / matrix
	x <- as.matrix(x)
	dnl <- length(dimnames(x)[[2]])
	nrx <- nrow(x)
	write.table(data.frame("## pre-customization"), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	write.table(data.frame(paste("\nnumber of columns in ",dnl," rows ",nrx,sep="")), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	# check
	isn <- sum(is.na(x))
	write.table(data.frame(paste("\nis NA number         ",isn,sep="")), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	iszero <- sum(x==0)
	write.table(data.frame(paste("\nis ZERO number       ",iszero," / ",dnl*nrx,sep="")), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	# remove all '0' rows
	a <- apply(x,1,sum)
	x <- x[!(a==0),]
	nrx <- nrow(x)
	write.table(data.frame(paste("\nremaining rows       ",nrx,sep="")), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	# adjust
	x[x<1 & x>=0] <- as.integer(1)
	cadj <- sum(x<1 & x>=0)
	write.table(data.frame(paste("\nremaining <1 ?       ",cadj,sep="")), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	mode(x) <- "integer"
	write.table(data.frame(paste("\nstorage.mode of x ?  ",storage.mode(x),sep="")), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	rxx <- range(x)
	write.table(data.frame(paste("\nrange of x           ",rxx[1]," - ",rxx[2],sep="")), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	write.table(data.frame(""), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	return(x)
}


# DE 1
de.fn <- function(countdata, a1, wpath, wfile, sfile, anno=NULL, sel=NULL, sign=F, test="DESeq", ret=F){
	# DESeq function
	# countdata: data matrix, samples in columns
	# a1: column data - conditions, vector e.g. c("A","N","N","N","A","N","A","C", ...)
	# wpath, wfile: path & file name results,  wpath, sfile: stat file result (for de.wrap() function)
	# anno: annotation: e.g. miR names - needs the row names of the countdata file
	# sel: the pairwise condition which should become a result: contrast: e.g. c("condition","A","N")  fc=A/N
	# test: standard: DESeq, other: DESeqMean, nbW, LRT see below
	# ret: return results or not
	coldata <- data.frame(		# numbers with as.factor()
			row.names=dimnames(countdata)[[2]],
			condition=as.factor(a1),
			type=as.factor(a1)
	)
	cat("\n")
	#print(coldata)
	cat("\n nrow coldata ",nrow(coldata))
	
	## QC on filtered data
	cat("\n countdata<0  ",sum(countdata[countdata<0]))
	cat("\n countdata==0 ",sum(countdata[countdata==0]))
#	return(list(countdata=countdata,coldata=coldata))
	
	## create deseq data structure for differential analysis
	# countdata:  rows: transcript IDs  cols: sample names
	dds.mat <- DESeqDataSetFromMatrix(
			countData = countdata,
			colData = coldata,
			#		design = ~ batch + condition)		# err -> Please read the vignette section 'Model matrix not full rank':  seems to be no work around
			design = ~ condition)
	
	dds.mat <- estimateSizeFactors(dds.mat)
	
	# check properties
	cat("\n nrow dds.mat ",nrow(dds.mat),"\n")
	
	## differential
	if(test=="DESeq"){
		dds.mat <- DESeq(dds.mat)
	}else if(test=="DESeqMean"){
		dds.mat <- DESeq(dds.mat, fitType="mean")		# fitType mean for sparse miR data -- sometimes even this is crashing
	}else if(test=="nbW"){
		dds.mat <- estimateDispersionsGeneEst(dds.mat)
		dispersions(dds.mat) <- mcols(dds.mat)$dispGeneEst
		dds.mat <- nbinomWaldTest(dds.mat)
	}else if(test=="LRT"){
		dds.mat <- estimateDispersionsGeneEst(dds.mat)
		dispersions(dds.mat) <- mcols(dds.mat)$dispGeneEst
		dds.mat <- nbinomLRT(dds.mat, reduced= ~1)
	}
	
	## results 1
	# list names
	if(is.null(sel)){
		cat("\n dds.mat obj ",resultsNames(dds.mat),"\n")
		cat("\n >>> choose a contrast argument and restart\n")
		return()
	}else{
		cat("\n chosen ",sel,"\n")
	}
	
	# A-B
	res.A.N <- results(dds.mat, contrast=sel)
	# order by adjusted p-value
	res.A.N <- as.data.frame(res.A.N)
	res.A.N[is.na(res.A.N$pvalue), "pvalue"] <- 1	# technically substitute NA by 1 for intersection
	res.A.N[is.na(res.A.N$padj), "padj"] <- 1	# technically substitute NA by 1 for intersection
	res.A.N <- res.A.N[order(res.A.N$pvalue), ]
	
	# number of significants
	yall <- nrow(res.A.N)
	y1true <- sum(res.A.N$pvalue<0.05)
	y2true <- sum(res.A.N$padj<0.05)
	cat("\n <0.05 :  FALSE  TRUE")
	cat("\n sign. p      ",yall-y1true," ",y1true)
	cat("\n sign. padj   ",yall-y2true," ",y2true)
	# write result numbers in a log file
	write.table(data.frame(pF=yall-y1true, pT=y1true, padjF=yall-y2true, padjT=y2true, row.names=wfile),
			file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=T, col.names=NA)
	write.table(data.frame(""), file=paste(wpath,sfile,".tsv",sep=""), append=T, sep="\t", dec=".", quote=F, row.names=F, col.names=F)
	
	# add additional annotations - needs row.names
	if(!is.null(anno)){
		res.A.N.df <- merge(as.data.frame(res.A.N), as.data.frame(anno), by="row.names", sort=F)
		row.names(res.A.N.df) <- res.A.N.df[,1]
		res.A.N.df <- res.A.N.df[,-1]
		res.A.N <- res.A.N.df
	}
	
	# create a  data+results  dataframe
	dcm <- as.data.frame(counts(dds.mat, normalized=T)-1)
	dcm <- round(dcm, 0)
	sti <- as.data.frame(res.A.N)
	#sti <- round(sti, 2)
	res.A.N.df <- merge(sti, dcm, by="row.names", sort=F)
	row.names(res.A.N.df) <- res.A.N.df[,1]
	res.A.N.df <- res.A.N.df[,-1]
	res.A.N.df <- res.A.N.df[order(res.A.N.df$padj), ]	# it is - but to assure
	# keep only sign.
	if(sign){
		res.A.N.df <- res.A.N.df[res.A.N.df$padj<0.05, ]
	}
	
	# write data DE used/adjusted counts
	write.table(x=res.A.N.df, file=paste(wpath,wfile,".tsv",sep=""), sep="\t", dec=".", col.names=NA)		# NA means colnames with first blank
	cat("\n\n")
	if(ret){ return(res.A.N.df) }else{ return() }
}


# Extract result rows from log file
de.overview <- function(p, f, ret=F){
	# extract from wrap.de() log file the DE result lines
	# xp: path name, f: file name with extension, ret: additionally return table to workspace
	con <- file(paste(p,"/",f,sep=""), "r")
	out <- "pF\tpT\tpadjF\tpadjT"
	while( TRUE ){
		line <- readLines(con, n=1)
		if(length(line)==0){ break }
		gt <- grepl(pattern="^\\tpF\\tpT\\tpadjF\\tpadjT",x=line, perl=T)
		if( gt ){
			out <- rbind(out, readLines(con, n=1))
		}
	}
	close(con)
	write.table(out, paste(p,"/0res_",f,sep=""), append=F, quote=F, sep="", row.names=F, col.names=F)
	out <- read.table(paste(p,"/0res_",f,sep=""), header=T, sep="\t", quote="", stringsAsFactors=F)
	if(ret){ return(out) }
}
#a <- de.overview(p="res20240606redoumi/DESeqMean", f="de_overview.tsv", ret=T)





# DE2 - LRT batch
#alternate batch correction - from karolinska
# much more significat ones, nevertheless -only- padj as criteria useful
# cave: no fold change criteria possible here, because not pairwise
#glMDSPlot(d1, groups=groupnew1, folder="mds")
#
#dds <- DESeqDataSetFromMatrix(countData=counts2,
#		colData=groupnew1,
#		design= ~ Batch + hCG)
#
#dds_lrt <- DESeq(dds, test="LRT",reduced= ~ Batch)
#
#resultsNames(dds_lrt) # lists the coefficients
#res_LRT <- results(dds_lrt, name="hCG_Positive_vs_Negative")
#hist(res_LRT$pvalue)
#DE <- as.data.frame(res_LRT)
#
#DE$FDR <- p.adjust(DE$pvalue, method="hochberg", n=length(DE$pvalue))
#
#length(which(DE$padj < 0.05))

de.fn2 <- function(countdata, a1, b1, wpath, wfile, LRT.a=T, sel=NULL, sign=F, ret=F){
	# DESeq function - batch correction
	# countdata: data matrix, samples in columns
	# a1: column data - conditions, vector e.g. c("A","N","N","N","A","N","A","C", ...)
	# b1: batch condition
	# wpath, wfile: path & file name
	# sel: the pairwise condition which should become a result:
	#  LRT.a=F: contrast: e.g. c("condition","A","N")  fc=A/N, or LRT.a=T: name: "condition_N_vs_A"
	coldata <- data.frame(		# numbers with as.factor()
			row.names=dimnames(countdata)[[2]],
			condition=as.factor(a1),
			batch=as.factor(b1)
	)
	cat("\n")
	cat("\n nrow coldata ",nrow(coldata),"\n sel structure \n")
	print(coldata)
	cat("\n")
	
	## QC on filtered data
	cat("\n countdata<0  ",sum(countdata[countdata<0]))
	cat("\n countdata==0 ",sum(countdata[countdata==0]))
#	return(list(countdata=countdata,coldata=coldata))
	
	## create deseq data structure for differential analysis
	# countdata:  rows: transcript IDs  cols: sample names
	dds.mat <- DESeqDataSetFromMatrix(
			countData = countdata,
			colData = coldata,
			design = ~ batch + condition)		# err -> Please read the vignette section 'Model matrix not full rank':  seems to be no work around
	
	dds.mat <- estimateSizeFactors(dds.mat)
	
	# check properties
	cat("\n nrow dds.mat ",nrow(dds.mat),"\n")
	
	## differential
	if(LRT.a){
		dds.lrt <- DESeq(dds.mat, test="LRT",reduced= ~ batch)
		# list names
		if(is.null(sel)){
			cat("\n dds.lrt obj ",resultsNames(dds.lrt),"\n")
			cat("\n choose a name argument \n")
			return()
		}else{
			cat("\n chosen ",sel,"\n")
		}
		
		res.LRT <- results(dds.lrt, name=sel) # one of resultnames
		DE <- as.data.frame(res.LRT)
		DE$padj <- p.adjust(DE$pvalue, method="hochberg", n=length(DE$pvalue))
		
		DE[is.na(DE$padj), "padj"] <- 1	# technically substitute NA by 1 for intersection
		DE <- DE[order(DE$pvalue), ]
		
		# number of significants
		cat("\n sign. padj    ",table(DE$padj<0.05))
		
		# create a  data+results  dataframe
		DE.df <- merge(as.data.frame(DE), as.data.frame(counts(dds.lrt, normalized=T)-1), by="row.names", sort=F)
		row.names(DE.df) <- DE.df[,1]
		DE.df <- DE.df[,-1]
		DE.df <- DE.df[order(DE.df$padj), ]	# it is - but to assure
	}else{
		dds.mat <- DESeq(dds.mat)
		# list names
		if(is.null(sel)){
			cat("\n dds.mat obj ",resultsNames(dds.mat),"\n")
			cat("\n choose a contrast argument \n")
			return()
		}else{
			cat("\n chosen ",sel,"\n")
		}
		
		# A-B
		DE <- results(dds.mat, contrast=sel)	# factor_name,char_numerator,char_denominator
		# order by adjusted p-value
		DE <- as.data.frame(DE)
		DE[is.na(DE$padj), "padj"] <- 1	# technically substitute NA by 1 for intersection
		DE <- DE[order(DE$pvalue), ]
		
		# number of significants
		cat("\n <0.05 : FALSE TRUE")
		cat("\n sign. padj    ",table(DE$padj<0.05))
		cat("\n sign. p       ",table(DE$pvalue<0.05))
		
		# create a  data+results  dataframe
		DE.df <- merge(as.data.frame(DE), as.data.frame(counts(dds.mat, normalized=T)-1), by="row.names", sort=F)
		row.names(DE.df) <- DE.df[,1]
		DE.df <- DE.df[,-1]
		DE.df <- DE.df[order(DE.df$padj), ]	# it is - but to assure
	}
	# keep only sign.
	if(sign){
		DE.df <- res.A.N.df[DE.df$padj<0.05, ]
	}
	
		
	# write data DE used/adjusted counts
	write.table(x=DE.df, file=paste(wpath,wfile,".tsv",sep=""), sep="\t", dec=".", col.names=NA)		# NA means colnames with first blank
	cat("\n\n")
	if(ret){ return(DE.df) }else{ return() }
}






