import Erdos1135.Terras.Parity.Word
import Mathlib.Data.Fintype.Powerset
import Mathlib.Data.Nat.Choose.Sum

/-!
# Terras/Everett Bitstring Counting

This module proves the first finite combinatorics theorem needed for the Terras/Everett density
route: parity words with exactly `r` odd bits are counted by `Nat.choose k r`.
-/

namespace Erdos1135
namespace Terras

/-- The support of odd bits in a parity word. -/
def parityWordSupport {k : ℕ} (w : ParityWord k) : Finset (Fin k) :=
  Finset.univ.filter fun i => w i

noncomputable def parityWordsWithNumOddEquivPowersetCard (k r : ℕ) :
    {w : ParityWord k // numOdd w = r} ≃
      {s : Finset (Fin k) // s ∈ (Finset.univ : Finset (Fin k)).powersetCard r} where
  toFun w := by
    refine ⟨parityWordSupport w.1, ?_⟩
    rw [Finset.mem_powersetCard]
    exact ⟨Finset.subset_univ _, by simpa [parityWordSupport, numOdd] using w.2⟩
  invFun s := by
    refine ⟨(fun i : Fin k => decide (i ∈ s.1)), ?_⟩
    have hs : s.1.card = r := (Finset.mem_powersetCard.mp s.2).2
    have hfilter : (Finset.univ.filter fun i : Fin k => decide (i ∈ s.1)).card = s.1.card := by
      congr
      ext i
      simp
    rw [numOdd, hfilter, hs]
  left_inv w := by
    ext i
    simp [parityWordSupport]
  right_inv s := by
    ext i
    simp [parityWordSupport]

theorem card_words_with_numOdd (k r : ℕ) :
    Fintype.card {w : ParityWord k // numOdd w = r} = Nat.choose k r := by
  classical
  rw [Fintype.card_congr (parityWordsWithNumOddEquivPowersetCard k r)]
  change Fintype.card ↑((Finset.univ : Finset (Fin k)).powersetCard r) = Nat.choose k r
  rw [Fintype.card_coe, Finset.card_powersetCard, Finset.card_univ, Fintype.card_fin]

lemma numOdd_le {k : ℕ} (w : ParityWord k) : numOdd w ≤ k := by
  rw [numOdd]
  exact (Finset.card_filter_le _ _).trans_eq (Fintype.card_fin k)

/-- Odd-count values that make the accelerated branch multiplier smaller than one. -/
def contractingOddCounts (k : ℕ) : Finset ℕ :=
  (Finset.range (k + 1)).filter fun r => 3 ^ r < 2 ^ k

lemma mem_contractingOddCounts_iff {k r : ℕ} :
    r ∈ contractingOddCounts k ↔ r ≤ k ∧ 3 ^ r < 2 ^ k := by
  simp [contractingOddCounts]

lemma numOdd_mem_range_succ {k : ℕ} (w : ParityWord k) :
    numOdd w ∈ Finset.range (k + 1) := by
  rw [Finset.mem_range, Nat.lt_succ_iff]
  exact numOdd_le w

lemma numOdd_mem_contractingOddCounts_iff {k : ℕ} (w : ParityWord k) :
    numOdd w ∈ contractingOddCounts k ↔ 3 ^ numOdd w < 2 ^ k := by
  rw [mem_contractingOddCounts_iff]
  exact ⟨fun h => h.2, fun h => ⟨numOdd_le w, h⟩⟩

theorem card_words_with_numOdd_and_contracting (k r : ℕ) :
    Fintype.card {w : ParityWord k // numOdd w = r ∧ 3 ^ numOdd w < 2 ^ k} =
      if 3 ^ r < 2 ^ k then Nat.choose k r else 0 := by
  classical
  by_cases h : 3 ^ r < 2 ^ k
  · rw [if_pos h]
    let e : {w : ParityWord k // numOdd w = r ∧ 3 ^ numOdd w < 2 ^ k} ≃
        {w : ParityWord k // numOdd w = r} :=
      { toFun := fun w => ⟨w.1, w.2.1⟩
        invFun := fun w => ⟨w.1, w.2, by simpa [w.2] using h⟩
        left_inv := by
          intro w
          ext i
          rfl
        right_inv := by
          intro w
          ext i
          rfl }
    rw [Fintype.card_congr e, card_words_with_numOdd]
  · rw [if_neg h]
    have hempty : IsEmpty {w : ParityWord k // numOdd w = r ∧ 3 ^ numOdd w < 2 ^ k} := by
      refine ⟨fun w => ?_⟩
      exact h (by simpa [w.2.1] using w.2.2)
    exact Fintype.card_eq_zero

noncomputable def contractingWordsEquivSigmaOddCount (k : ℕ) :
    {w : ParityWord k // 3 ^ numOdd w < 2 ^ k} ≃
      Sigma (fun r : {r : ℕ // r ∈ contractingOddCounts k} =>
        {w : ParityWord k // numOdd w = r.1}) where
  toFun w :=
    ⟨⟨numOdd w.1, (numOdd_mem_contractingOddCounts_iff w.1).mpr w.2⟩, ⟨w.1, rfl⟩⟩
  invFun rw :=
    ⟨rw.2.1, by
      rw [rw.2.2]
      exact (mem_contractingOddCounts_iff.mp rw.1.2).2⟩
  left_inv w := by
    ext i
    rfl
  right_inv rw := by
    cases rw with
    | mk r w =>
      cases r with
      | mk r hr =>
        cases w with
        | mk w hw =>
          simp at hw ⊢
          constructor
          · exact hw
          · subst r
            simp

theorem card_contracting_words_eq_sum_choose (k : ℕ) :
    Fintype.card {w : ParityWord k // 3 ^ numOdd w < 2 ^ k} =
      ∑ r ∈ contractingOddCounts k, Nat.choose k r := by
  classical
  rw [Fintype.card_congr (contractingWordsEquivSigmaOddCount k)]
  rw [Fintype.card_sigma]
  simp [card_words_with_numOdd]
  rw [Finset.sum_attach]

end Terras
end Erdos1135
