import Erdos1135.KrasikovLagarias.AdaptiveAllowedFamily
import Erdos1135.KrasikovLagarias.AdaptivePathExtraction
import Erdos1135.KrasikovLagarias.EliminationNormalizer

/-!
# Finite normal forms for the adaptive KL policy
-/

namespace Erdos1135
namespace KrasikovLagarias
namespace AdaptiveEliminationTree

open AdaptiveEliminationAllowedFamily
open AdaptiveEliminationPolicy
open EliminationCriticalTree
open EliminationPolicy
open EliminationResidue

mutual

inductive Normalized {k : Nat} {hk : 2 ≤ k}
    (potential : AdaptiveEliminationPolicy.ForcedPotential k hk) :
    State k → Tree k → Prop where
  | terminal (state : State k) (hretarded : state.current.value < 0) :
      Normalized potential state (.terminal state.current)
  | l1 (state : State k) (hadvanced : 0 ≤ state.current.value)
      (hrow : rowKind state.current.index = FiniteCertificate.RowKind.l1)
      {principal : Tree k} {auxiliary : NonemptyFamily k}
      (principalNormalized : Normalized potential
        (state.descend (state.current.principalChild
          (Nat.le_trans (by omega) hk))) principal)
      (auxiliaryNormalized : NormalizedFamily potential state
        (d1Family potential state hrow) auxiliary) :
      Normalized potential state
        (.principalMin state.current principal auxiliary)
  | l2 (state : State k) (hadvanced : 0 ≤ state.current.value)
      (hrow : rowKind state.current.index = FiniteCertificate.RowKind.l2)
      {principal : Tree k}
      (principalNormalized : Normalized potential
        (state.descend (state.current.principalChild
          (Nat.le_trans (by omega) hk))) principal) :
      Normalized potential state (.principalOnly state.current principal)
  | l3 (state : State k) (hadvanced : 0 ≤ state.current.value)
      (hrow : rowKind state.current.index = FiniteCertificate.RowKind.l3)
      {principal : Tree k} {auxiliary : NonemptyFamily k}
      (principalNormalized : Normalized potential
        (state.descend (state.current.principalChild
          (Nat.le_trans (by omega) hk))) principal)
      (auxiliaryNormalized : NormalizedFamily potential state
        (d3Family potential state hrow) auxiliary) :
      Normalized potential state
        (.principalMin state.current principal auxiliary)

inductive NormalizedFamily {k : Nat} {hk : 2 ≤ k}
    (potential : AdaptiveEliminationPolicy.ForcedPotential k hk) :
    State k → NonemptyFamily k → NonemptyFamily k → Prop where
  | one {parent : State k} (label : Label k) (tree : Tree k)
      (hlegal : AdaptiveEliminationPolicy.LegalChild potential
        (parent.descend label) parent)
      (normalized : Normalized potential (parent.descend label) tree) :
      NormalizedFamily potential parent (.one (.terminal label)) (.one tree)
  | cons {parent : State k} (label : Label k) (tree : Tree k)
      {sourceTail normalizedTail : NonemptyFamily k}
      (hlegal : AdaptiveEliminationPolicy.LegalChild potential
        (parent.descend label) parent)
      (normalized : Normalized potential (parent.descend label) tree)
      (tailNormalized : NormalizedFamily potential parent
        sourceTail normalizedTail) :
      NormalizedFamily potential parent
        (.cons (.terminal label) sourceTail) (.cons tree normalizedTail)

end

namespace Normalized

theorem rootLabel_eq {k : Nat} {hk : 2 ≤ k}
    {potential : AdaptiveEliminationPolicy.ForcedPotential k hk}
    {state : State k} {tree : Tree k}
    (normalized : Normalized potential state tree) :
    tree.rootLabel = state.current := by
  cases normalized <;> rfl

end Normalized

mutual

theorem Normalized.terminalsRetarded {k : Nat} {hk : 2 ≤ k}
    {potential : AdaptiveEliminationPolicy.ForcedPotential k hk}
    {state : State k} {tree : Tree k}
    (normalized : Normalized potential state tree) :
    tree.TerminalsRetarded := by
  cases normalized with
  | terminal state hretarded => exact hretarded
  | l1 state hadvanced hrow principalNormalized auxiliaryNormalized =>
      exact ⟨principalNormalized.terminalsRetarded,
        auxiliaryNormalized.terminalsRetarded⟩
  | l2 state hadvanced hrow principalNormalized =>
      exact principalNormalized.terminalsRetarded
  | l3 state hadvanced hrow principalNormalized auxiliaryNormalized =>
      exact ⟨principalNormalized.terminalsRetarded,
        auxiliaryNormalized.terminalsRetarded⟩

theorem NormalizedFamily.terminalsRetarded {k : Nat} {hk : 2 ≤ k}
    {potential : AdaptiveEliminationPolicy.ForcedPotential k hk}
    {parent : State k} {source normalized : NonemptyFamily k}
    (familyNormalized : NormalizedFamily potential parent source normalized) :
    normalized.TerminalsRetarded := by
  cases familyNormalized with
  | one label tree hlegal normalized => exact normalized.terminalsRetarded
  | cons label tree hlegal normalized tailNormalized =>
      exact ⟨normalized.terminalsRetarded, tailNormalized.terminalsRetarded⟩

end


private theorem d1_member_legal {k : Nat} {hk : 2 ≤ k}
    (potential : AdaptiveEliminationPolicy.ForcedPotential k hk)
    (parent : State k) (hadvanced : 0 ≤ parent.current.value)
    (hrow : rowKind parent.current.index = FiniteCertificate.RowKind.l1)
    {source : Tree k}
    (hmember : NonemptyFamily.Member source (d1Family potential parent hrow)) :
    ∃ label : Label k, source = .terminal label ∧
      AdaptiveEliminationPolicy.LegalChild potential
        (parent.descend label) parent := by
  have hmember' : NonemptyFamily.Member source
      (allowedFamily parent.history (parent.current.d1Child hk hrow)
        (potential.d1Lift parent.current.index hrow)) := by
    simpa [d1Family] using hmember
  obtain ⟨lift, rfl⟩ := member_allowedFamily_exists parent.history
    (parent.current.d1Child hk hrow)
    (potential.d1Lift parent.current.index hrow) hmember'
  exact ⟨_, rfl, d1_legalChild_of_terminal_member potential parent
    hadvanced hrow lift hmember⟩

private theorem d3_member_legal {k : Nat} {hk : 2 ≤ k}
    (potential : AdaptiveEliminationPolicy.ForcedPotential k hk)
    (parent : State k) (hadvanced : 0 ≤ parent.current.value)
    (hrow : rowKind parent.current.index = FiniteCertificate.RowKind.l3)
    {source : Tree k}
    (hmember : NonemptyFamily.Member source (d3Family potential parent hrow)) :
    ∃ label : Label k, source = .terminal label ∧
      AdaptiveEliminationPolicy.LegalChild potential
        (parent.descend label) parent := by
  have hmember' : NonemptyFamily.Member source
      (allowedFamily parent.history (parent.current.d3Child hk hrow)
        (potential.d3Lift parent.current.index hrow)) := by
    simpa [d3Family] using hmember
  obtain ⟨lift, rfl⟩ := member_allowedFamily_exists parent.history
    (parent.current.d3Child hk hrow)
    (potential.d3Lift parent.current.index hrow) hmember'
  exact ⟨_, rfl, d3_legalChild_of_terminal_member potential parent
    hadvanced hrow lift hmember⟩

private theorem normalizedFamily_exists_of_legal {k : Nat} {hk : 2 ≤ k}
    (potential : AdaptiveEliminationPolicy.ForcedPotential k hk)
    (parent : State k) (source : NonemptyFamily k)
    (hlegal : ∀ sourceTree,
      NonemptyFamily.Member sourceTree source →
        ∃ label : Label k, sourceTree = .terminal label ∧
          AdaptiveEliminationPolicy.LegalChild potential
            (parent.descend label) parent)
    (normalizeChild : ∀ child,
      AdaptiveEliminationPolicy.LegalChild potential child parent →
        ∃ tree, Normalized potential child tree) :
    ∃ normalized, NormalizedFamily potential parent source normalized := by
  cases source with
  | one sourceTree =>
      obtain ⟨label, rfl, edge⟩ := hlegal sourceTree rfl
      obtain ⟨tree, treeNormalized⟩ := normalizeChild _ edge
      exact ⟨.one tree, NormalizedFamily.one label tree edge treeNormalized⟩
  | cons sourceTree sourceTail =>
      obtain ⟨label, rfl, edge⟩ := hlegal sourceTree (Or.inl rfl)
      obtain ⟨tree, treeNormalized⟩ := normalizeChild _ edge
      obtain ⟨normalizedTail, tailNormalized⟩ :=
        normalizedFamily_exists_of_legal potential parent sourceTail
          (fun child hmember => hlegal child (Or.inr hmember)) normalizeChild
      exact ⟨.cons tree normalizedTail,
        NormalizedFamily.cons label tree edge treeNormalized tailNormalized⟩
termination_by source

theorem normalized_exists_of_acc {k : Nat} {hk : 2 ≤ k}
    (potential : AdaptiveEliminationPolicy.ForcedPotential k hk)
    (state : State k)
    (hacc : Acc (AdaptiveEliminationPolicy.LegalChild potential) state) :
    ∃ tree, Normalized potential state tree := by
  induction hacc with
  | intro state childrenAccessible normalizeChild =>
      by_cases hretarded : state.current.value < 0
      · exact ⟨.terminal state.current, Normalized.terminal state hretarded⟩
      · have hadvanced : 0 ≤ state.current.value := le_of_not_gt hretarded
        have principalEdge : AdaptiveEliminationPolicy.LegalChild potential
            (state.descend (state.current.principalChild
              (Nat.le_trans (by omega) hk))) state :=
          AdaptiveEliminationPolicy.LegalChild.principal state hadvanced
        obtain ⟨principal, principalNormalized⟩ :=
          normalizeChild _ principalEdge
        cases hrow : rowKind state.current.index with
        | l1 =>
            obtain ⟨auxiliary, auxiliaryNormalized⟩ :=
              normalizedFamily_exists_of_legal potential state
                (d1Family potential state hrow)
                (fun source hmember =>
                  d1_member_legal potential state hadvanced hrow hmember)
                normalizeChild
            exact ⟨.principalMin state.current principal auxiliary,
              Normalized.l1 state hadvanced hrow principalNormalized
                auxiliaryNormalized⟩
        | l2 =>
            exact ⟨.principalOnly state.current principal,
              Normalized.l2 state hadvanced hrow principalNormalized⟩
        | l3 =>
            obtain ⟨auxiliary, auxiliaryNormalized⟩ :=
              normalizedFamily_exists_of_legal potential state
                (d3Family potential state hrow)
                (fun source hmember =>
                  d3_member_legal potential state hadvanced hrow hmember)
                normalizeChild
            exact ⟨.principalMin state.current principal auxiliary,
              Normalized.l3 state hadvanced hrow principalNormalized
                auxiliaryNormalized⟩

theorem root_normalized_exists {k : Nat} {hk : 2 ≤ k}
    (potential : AdaptiveEliminationPolicy.ForcedPotential k hk)
    (rootIndex : PrincipalIndex k) :
    ∃ tree, Normalized potential (State.root rootIndex) tree :=
  normalized_exists_of_acc potential (State.root rootIndex)
    (AdaptiveEliminationPolicy.root_accessible potential rootIndex)

end AdaptiveEliminationTree
end KrasikovLagarias
end Erdos1135
