import Erdos1135.KrasikovLagarias.AdaptiveForcedPotential
import Erdos1135.KrasikovLagarias.ForcedPotentialCertificate

/-!
# Finite certificates for adaptive KL forced-transition potentials

The payload stores only the bounded residue potential.  On each D1 or D3 row
the checker deterministically selects the least-potential lift (breaking ties
in the order `0, 1, 2`) and verifies the selected-lift inequality required by
`AdaptiveEliminationPolicy.ForcedPotential`.
-/

namespace Erdos1135
namespace KrasikovLagarias

open EliminationResidue

/-- Finite byte-valued data for an adaptive selected-lift potential. -/
structure AdaptiveForcedPotentialCertificate where
  k : Nat
  bound : Nat
  values : Array Nat
  deriving DecidableEq, Repr

namespace AdaptiveForcedPotentialCertificate

def MetadataValid (cert : AdaptiveForcedPotentialCertificate) : Prop :=
  2 ≤ cert.k ∧
    cert.values.size = principalCount cert.k ∧
    ∀ index : Fin cert.values.size, cert.values[index] ≤ cert.bound

instance metadataValidDecidable (cert : AdaptiveForcedPotentialCertificate) :
    Decidable cert.MetadataValid := by
  unfold MetadataValid
  infer_instance

def metadataCheck (cert : AdaptiveForcedPotentialCertificate) : Bool :=
  decide (2 ≤ cert.k) &&
    decide (cert.values.size = principalCount cert.k) &&
      cert.values.all fun entry => decide (entry ≤ cert.bound)

theorem metadataCheck_eq_true_iff (cert : AdaptiveForcedPotentialCertificate) :
    cert.metadataCheck = true ↔ cert.MetadataValid := by
  simp only [metadataCheck, Bool.and_eq_true, decide_eq_true_eq,
    Array.all_eq_true]
  constructor
  · rintro ⟨⟨hk, hsize⟩, hall⟩
    refine ⟨hk, hsize, ?_⟩
    intro index
    exact hall index.val index.isLt
  · rintro ⟨hk, hsize, hall⟩
    refine ⟨⟨hk, hsize⟩, ?_⟩
    intro index hindex
    exact hall ⟨index, hindex⟩

def value (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) : Nat :=
  cert.values[index.val]'(by
    rw [hmeta.2.1]
    exact index.isLt)

theorem value_le_bound (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) :
    cert.value hmeta index ≤ cert.bound := by
  exact hmeta.2.2 ⟨index.val, by
    rw [hmeta.2.1]
    exact index.isLt⟩

/-- Deterministic argmin of three natural values, with least-index tie
breaking. -/
def leastLift (v0 v1 v2 : Nat) : LiftIndex :=
  if v0 ≤ v1 ∧ v0 ≤ v2 then 0
  else if v1 ≤ v2 then 1
  else 2

def d1Lift (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k)
    (hrow : rowKind index = FiniteCertificate.RowKind.l1) : LiftIndex :=
  leastLift
    (cert.value hmeta (d1PrincipalLift hmeta.1 index hrow 0))
    (cert.value hmeta (d1PrincipalLift hmeta.1 index hrow 1))
    (cert.value hmeta (d1PrincipalLift hmeta.1 index hrow 2))

def d3Lift (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k)
    (hrow : rowKind index = FiniteCertificate.RowKind.l3) : LiftIndex :=
  leastLift
    (cert.value hmeta (d3PrincipalLift hmeta.1 index hrow 0))
    (cert.value hmeta (d3PrincipalLift hmeta.1 index hrow 1))
    (cert.value hmeta (d3PrincipalLift hmeta.1 index hrow 2))

def PrincipalValid (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) : Prop :=
  cert.value hmeta
      (fourIndex (Nat.le_trans (by omega) hmeta.1) index) ≤
    cert.value hmeta index + 6

def D1Valid (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) : Prop :=
  if hrow : rowKind index = FiniteCertificate.RowKind.l1 then
    cert.value hmeta
        (d1PrincipalLift hmeta.1 index hrow
          (cert.d1Lift hmeta index hrow)) ≤
      cert.value hmeta index + 1
  else
    True

def D3Valid (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) : Prop :=
  if hrow : rowKind index = FiniteCertificate.RowKind.l3 then
    cert.value hmeta
        (d3PrincipalLift hmeta.1 index hrow
          (cert.d3Lift hmeta index hrow)) + 2 ≤
      cert.value hmeta index
  else
    True

instance principalValidDecidable (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) :
    Decidable (cert.PrincipalValid hmeta index) := by
  unfold PrincipalValid
  infer_instance

instance d1ValidDecidable (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) :
    Decidable (cert.D1Valid hmeta index) := by
  unfold D1Valid
  split <;> infer_instance

instance d3ValidDecidable (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) :
    Decidable (cert.D3Valid hmeta index) := by
  unfold D3Valid
  split <;> infer_instance

def RowValid (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) : Prop :=
  cert.PrincipalValid hmeta index ∧
    cert.D1Valid hmeta index ∧ cert.D3Valid hmeta index

instance rowValidDecidable (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) (index : PrincipalIndex cert.k) :
    Decidable (cert.RowValid hmeta index) := by
  unfold RowValid
  infer_instance

structure Valid (cert : AdaptiveForcedPotentialCertificate) : Prop where
  metadata : cert.MetadataValid
  rows : ∀ index : PrincipalIndex cert.k, cert.RowValid metadata index

def rowsCheck (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) : Bool :=
  (List.finRange (principalCount cert.k)).all fun index =>
    decide (cert.RowValid hmeta index)

theorem rowsCheck_eq_true_iff (cert : AdaptiveForcedPotentialCertificate)
    (hmeta : cert.MetadataValid) :
    cert.rowsCheck hmeta = true ↔
      ∀ index : PrincipalIndex cert.k, cert.RowValid hmeta index := by
  rw [rowsCheck, List.all_eq_true]
  constructor
  · intro hall index
    exact of_decide_eq_true (hall index (List.mem_finRange index))
  · intro hall index hindex
    exact decide_eq_true (hall index)

def check (cert : AdaptiveForcedPotentialCertificate) : Bool :=
  if hcheck : cert.metadataCheck = true then
    cert.rowsCheck ((cert.metadataCheck_eq_true_iff).mp hcheck)
  else
    false

theorem check_eq_true_iff (cert : AdaptiveForcedPotentialCertificate) :
    cert.check = true ↔ cert.Valid := by
  by_cases hcheck : cert.metadataCheck = true
  · constructor
    · intro hvalid
      rw [check, dif_pos hcheck] at hvalid
      exact ⟨(cert.metadataCheck_eq_true_iff).mp hcheck,
        (cert.rowsCheck_eq_true_iff _).mp hvalid⟩
    · intro hvalid
      rw [check, dif_pos hcheck]
      apply (cert.rowsCheck_eq_true_iff _).mpr
      simpa only [Subsingleton.elim hvalid.metadata
        ((cert.metadataCheck_eq_true_iff).mp hcheck)] using hvalid.rows
  · constructor
    · intro hvalid
      simp [check, hcheck] at hvalid
    · intro hvalid
      exact (hcheck ((cert.metadataCheck_eq_true_iff).mpr hvalid.metadata)).elim

theorem valid_of_check {cert : AdaptiveForcedPotentialCertificate}
    (hcheck : cert.check = true) : cert.Valid :=
  (cert.check_eq_true_iff).mp hcheck

def Valid.toForcedPotential {cert : AdaptiveForcedPotentialCertificate}
    (hvalid : cert.Valid) :
    AdaptiveEliminationPolicy.ForcedPotential cert.k hvalid.metadata.1 where
  value := cert.value hvalid.metadata
  bound := cert.bound
  value_le_bound := cert.value_le_bound hvalid.metadata
  principal_le := fun index => (hvalid.rows index).1
  d1Lift := cert.d1Lift hvalid.metadata
  d1Selected_le := by
    intro index hrow
    simpa [D1Valid, hrow] using (hvalid.rows index).2.1
  d3Lift := cert.d3Lift hvalid.metadata
  d3Selected_le := by
    intro index hrow
    simpa [D3Valid, hrow] using (hvalid.rows index).2.2

def forcedPotential_of_check {cert : AdaptiveForcedPotentialCertificate}
    (hcheck : cert.check = true) :
    AdaptiveEliminationPolicy.ForcedPotential cert.k
      (cert.valid_of_check hcheck).metadata.1 :=
  (cert.valid_of_check hcheck).toForcedPotential

end AdaptiveForcedPotentialCertificate

def adaptivePotentialK2Canary : AdaptiveForcedPotentialCertificate where
  k := 2
  bound := 2
  values := decodeU8Hex "000002"

theorem adaptivePotentialK2Canary_check :
    adaptivePotentialK2Canary.check = true := by
  native_decide

end KrasikovLagarias
end Erdos1135
