import Mathlib.Analysis.SpecialFunctions.Log.Basic
import Mathlib.Data.List.Sort
import Mathlib.Tactic

/-!
A mathematical translation of DESeq2 1.52.0, core.R:536–579, ratio branch.
The theorem is over real numbers, not R binary64 arithmetic.
Rows supplied to this model are the nonempty retained positive rows.
-/
namespace DESeq2Verification
noncomputable section
open scoped BigOperators

/-- R's median: average the two central order statistics.
For odd length both indices select the same element. Empty is a model-only sentinel. -/
def median (xs : List ℝ) : ℝ :=
  if h : xs.length = 0 then 0 else
    let ys := xs.mergeSort (fun a b => decide (a ≤ b))
    (ys[(xs.length - 1) / 2]'(by have hy : ys.length = xs.length := List.length_mergeSort xs; omega) +
      ys[xs.length / 2]'(by have hy : ys.length = xs.length := List.length_mergeSort xs; omega)) / 2

/-- Translation commutes with sorting and both odd and even medians. -/
theorem median_add (xs : List ℝ) (hne : xs ≠ []) (d : ℝ) :
    median (xs.map (fun x => x + d)) = median xs + d := by
  have hn : xs.length ≠ 0 := by simpa using hne
  have hs : (xs.mergeSort (fun a b => decide (a ≤ b))).map (fun x => x + d) =
      (xs.map (fun x => x + d)).mergeSort (fun a b => decide (a ≤ b)) := by
    apply List.map_mergeSort
    intro a ha b hb
    simp only [add_le_add_iff_right]
  simp only [median, List.length_map, hn, ↓reduceDIte]
  simp only [← hs, List.getElem_map]
  ring

/-- Mean across samples. -/
def meanLog {m : ℕ} (row : Fin m → ℝ) : ℝ :=
  (∑ k, Real.log (row k)) / m

def residuals {m : ℕ} (rows : List (Fin m → ℝ)) (j : Fin m) : List ℝ :=
  rows.map (fun row => Real.log (row j) - meanLog row)

/-- Exact-real translation of exp(median(log(count)-rowMeans(log(count)))). -/
def sizeFactor {m : ℕ} (rows : List (Fin m → ℝ)) (j : Fin m) : ℝ :=
  Real.exp (median (residuals rows j))

def scaleRows {m : ℕ} (rows : List (Fin m → ℝ)) (c : Fin m → ℝ) :=
  rows.map (fun row k => c k * row k)

def geometricMean {m : ℕ} (c : Fin m → ℝ) := Real.exp (meanLog c)

theorem meanLog_mul {m : ℕ} (c row : Fin m → ℝ)
    (hc : ∀ k, 0 < c k) (hr : ∀ k, 0 < row k) :
    meanLog (fun k => c k * row k) = meanLog c + meanLog row := by
  simp only [meanLog, Real.log_mul (ne_of_gt (hc _)) (ne_of_gt (hr _))]
  rw [Finset.sum_add_distrib, add_div]

theorem residuals_scale {m : ℕ} (rows : List (Fin m → ℝ)) (c : Fin m → ℝ)
    (hc : ∀ k, 0 < c k) (hr : ∀ row ∈ rows, ∀ k, 0 < row k) (j : Fin m) :
    residuals (scaleRows rows c) j =
      (residuals rows j).map (fun x => x + (Real.log (c j) - meanLog c)) := by
  simp only [residuals, scaleRows, List.map_map]
  apply List.map_congr_left
  intro row hrow
  dsimp
  rw [meanLog_mul c row hc (hr row hrow),
      Real.log_mul (ne_of_gt (hc j)) (ne_of_gt (hr row hrow j))]
  ring

/-- Main property: positive sample scaling changes size factors by c[j]/GM(c). -/
theorem sizeFactor_scale {m : ℕ} (rows : List (Fin m → ℝ))
    (hne : rows ≠ []) (c : Fin m → ℝ) (hc : ∀ k, 0 < c k)
    (hr : ∀ row ∈ rows, ∀ k, 0 < row k) (j : Fin m) :
    sizeFactor (scaleRows rows c) j = c j / geometricMean c * sizeFactor rows j := by
  have hn : residuals rows j ≠ [] := by simpa [residuals] using hne
  rw [sizeFactor, residuals_scale rows c hc hr j, median_add _ hn]
  rw [Real.exp_add, Real.exp_sub, Real.exp_log (hc j)]
  simp only [geometricMean, sizeFactor]
  ring

/-- Normalized counts change by the common factor GM(c), independent of sample. -/
theorem normalized_scaled {m : ℕ} (rows : List (Fin m → ℝ))
    (hne : rows ≠ []) (c : Fin m → ℝ) (hc : ∀ k, 0 < c k)
    (hr : ∀ row ∈ rows, ∀ k, 0 < row k) (j : Fin m) (x : ℝ) :
    (c j * x) / sizeFactor (scaleRows rows c) j =
      geometricMean c * (x / sizeFactor rows j) := by
  rw [sizeFactor_scale rows hne c hc hr j]
  have hsf : sizeFactor rows j ≠ 0 := Real.exp_ne_zero _
  have hgm : geometricMean c ≠ 0 := Real.exp_ne_zero _
  field_simp [ne_of_gt (hc j), hsf, hgm]

/-- Ratios of normalized sample values are invariant under positive sample scaling. -/
theorem normalized_ratio_scaled {m : ℕ} (rows : List (Fin m → ℝ))
    (hne : rows ≠ []) (c : Fin m → ℝ) (hc : ∀ k, 0 < c k)
    (hr : ∀ row ∈ rows, ∀ k, 0 < row k) (j k : Fin m) (x y : ℝ) :
    ((c j * x) / sizeFactor (scaleRows rows c) j) /
        ((c k * y) / sizeFactor (scaleRows rows c) k) =
      (x / sizeFactor rows j) / (y / sizeFactor rows k) := by
  rw [normalized_scaled rows hne c hc hr j x, normalized_scaled rows hne c hc hr k y]
  exact mul_div_mul_left _ _ (Real.exp_ne_zero _)

/-- Global depth scaling cancels from the estimated size factors. -/
theorem sizeFactor_global_scale {m : ℕ} (hm : 0 < m) (rows : List (Fin m → ℝ))
    (hne : rows ≠ []) (a : ℝ) (ha : 0 < a)
    (hr : ∀ row ∈ rows, ∀ k, 0 < row k) (j : Fin m) :
    sizeFactor (scaleRows rows (fun _ => a)) j = sizeFactor rows j := by
  rw [sizeFactor_scale rows hne (fun _ => a) (fun _ => ha) hr j]
  have hm' : (m : ℝ) ≠ 0 := by exact_mod_cast (Nat.ne_of_gt hm)
  have hg : geometricMean (fun _ : Fin m => a) = a := by
    simp [geometricMean, meanLog, hm', Real.exp_log ha]
  rw [hg, div_self (ne_of_gt ha), one_mul]

/-- Mathematical eligibility rule for the ratio branch on nonnegative finite counts.
R's -Inf log-geometric mean discards each row containing a zero. -/
def retainedRows {m : ℕ} (rows : List (Fin m → ℝ)) :=
  rows.filter (fun row => decide (∀ k, 0 < row k))

theorem retained_scale {m : ℕ} (rows : List (Fin m → ℝ)) (c : Fin m → ℝ)
    (hc : ∀ k, 0 < c k) :
    retainedRows (scaleRows rows c) = scaleRows (retainedRows rows) c := by
  simp only [retainedRows, scaleRows, List.filter_map]
  congr 1
  congr 1
  funext row
  simp only [Function.comp_apply, mul_pos_iff_of_pos_left (hc _)]

def estimateRatio {m : ℕ} (rows : List (Fin m → ℝ)) (j : Fin m) :=
  sizeFactor (retainedRows rows) j

/-- Scaling preserves row eligibility and equivariance, including discarded zero rows.
The nonempty retained-row hypothesis excludes the upstream all-genes-zero error. -/
theorem estimateRatio_scale {m : ℕ} (rows : List (Fin m → ℝ))
    (hne : retainedRows rows ≠ []) (c : Fin m → ℝ) (hc : ∀ k, 0 < c k) (j : Fin m) :
    estimateRatio (scaleRows rows c) j = c j / geometricMean c * estimateRatio rows j := by
  simp only [estimateRatio, retained_scale rows c hc]
  apply sizeFactor_scale _ hne c hc
  intro row hrow
  exact of_decide_eq_true (List.mem_filter.mp hrow).2

/-- Exact rational certificate predicate for recorded production outputs. -/
def WithinRelative (actual expected epsilon : ℚ) : Prop :=
  0 < expected ∧ 0 ≤ epsilon ∧ |actual - expected| ≤ epsilon * expected

instance (actual expected epsilon : ℚ) : Decidable (WithinRelative actual expected epsilon) :=
  inferInstanceAs (Decidable (0 < expected ∧ 0 ≤ epsilon ∧ |actual - expected| ≤ epsilon * expected))

def checkRelative (actual expected epsilon : ℚ) : Bool :=
  decide (WithinRelative actual expected epsilon)

/-- Checker acceptance entails the stated exact relative-error bound. -/
theorem checkRelative_sound (actual expected epsilon : ℚ)
    (h : checkRelative actual expected epsilon = true) :
    |actual - expected| / expected ≤ epsilon := by
  have hw : WithinRelative actual expected epsilon := of_decide_eq_true h
  exact (div_le_iff₀ hw.1).mpr hw.2.2

#print axioms median_add
#print axioms sizeFactor_scale
#print axioms normalized_scaled
#print axioms normalized_ratio_scaled
#print axioms sizeFactor_global_scale
#print axioms estimateRatio_scale
#print axioms checkRelative_sound
end
end DESeq2Verification
