{-# OPTIONS --safe --cubical #-}

module Spartan6.Netlist.StableDAG where

open import Spartan6.Prelude

open import Cubical.Data.FinData.Base using (weakenFin; fromℕ)
import Spartan6.Netlist.Checked as Checked
import Spartan6.Netlist.DAGEvaluation as Direct
import Spartan6.Primitive.LUT as LUT
import Spartan6.Semantics.Design as Semantics

-- Append-stable local indices use their ordinary numeric order: zero denotes
-- the oldest node and `fromℕ localCount` denotes a freshly appended node.

data StableWire (inputCount stateCount localCount : ℕ) : Type₀ where
  stableExternal : Fin inputCount
    → StableWire inputCount stateCount localCount
  stableStored : Fin stateCount
    → StableWire inputCount stateCount localCount
  stableLocal : Fin localCount
    → StableWire inputCount stateCount localCount
  stableLiteral : Bit
    → StableWire inputCount stateCount localCount

data StableNode (inputCount stateCount localCount : ℕ) : Type₀ where
  stableInvert : StableWire inputCount stateCount localCount
    → StableNode inputCount stateCount localCount
  stableAnd : StableWire inputCount stateCount localCount
    → StableWire inputCount stateCount localCount
    → StableNode inputCount stateCount localCount
  stableOr : StableWire inputCount stateCount localCount
    → StableWire inputCount stateCount localCount
    → StableNode inputCount stateCount localCount
  stableXor : StableWire inputCount stateCount localCount
    → StableWire inputCount stateCount localCount
    → StableNode inputCount stateCount localCount
  stableMux : StableWire inputCount stateCount localCount
    → StableWire inputCount stateCount localCount
    → StableWire inputCount stateCount localCount
    → StableNode inputCount stateCount localCount
  stableLUT : ∀ {arity}
    → LUT.TruthTable arity
    → Vec (StableWire inputCount stateCount localCount) arity
    → StableNode inputCount stateCount localCount

infixl 4 _▹_

data StableNodes (inputCount stateCount : ℕ) : ℕ → Type₀ where
  stableNoNodes : StableNodes inputCount stateCount 0
  _▹_ : ∀ {localCount}
    → StableNodes inputCount stateCount localCount
    → StableNode inputCount stateCount localCount
    → StableNodes inputCount stateCount (suc localCount)

snoc : ∀ {ℓ} {A : Type ℓ} {count}
  → Vec A count → A → Vec A (suc count)
snoc [] value = value ∷ []
snoc (item ∷ items) value = item ∷ snoc items value

reverseValues : ∀ {ℓ} {A : Type ℓ} {count}
  → Vec A count → Vec A count
reverseValues [] = []
reverseValues (item ∷ items) = snoc (reverseValues items) item

stableEvaluateWire : ∀ {inputCount stateCount localCount}
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
  → StableWire inputCount stateCount localCount
  → Bit
stableEvaluateWire external state locals (stableExternal index) =
  lookup index external
stableEvaluateWire external state locals (stableStored index) =
  lookup index state
stableEvaluateWire external state locals (stableLocal index) =
  lookup index locals
stableEvaluateWire external state locals (stableLiteral bit) = bit

stableEvaluateWires : ∀ {inputCount stateCount localCount count}
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
  → Vec (StableWire inputCount stateCount localCount) count
  → Vec Bit count
stableEvaluateWires external state locals [] = []
stableEvaluateWires external state locals (wire ∷ wires) =
  stableEvaluateWire external state locals wire
  ∷ stableEvaluateWires external state locals wires

stableEvaluateNode : ∀ {inputCount stateCount localCount}
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
  → StableNode inputCount stateCount localCount
  → Bit
stableEvaluateNode external state locals (stableInvert wire) =
  not (stableEvaluateWire external state locals wire)
stableEvaluateNode external state locals (stableAnd left right) =
  stableEvaluateWire external state locals left
  and stableEvaluateWire external state locals right
stableEvaluateNode external state locals (stableOr left right) =
  stableEvaluateWire external state locals left
  or stableEvaluateWire external state locals right
stableEvaluateNode external state locals (stableXor left right) =
  stableEvaluateWire external state locals left
  ⊕ stableEvaluateWire external state locals right
stableEvaluateNode external state locals
  (stableMux select when-low when-high) =
  mux (stableEvaluateWire external state locals select)
      (stableEvaluateWire external state locals when-low)
      (stableEvaluateWire external state locals when-high)
stableEvaluateNode external state locals (stableLUT table arguments) =
  LUT.evalLUT table
    (stableEvaluateWires external state locals arguments)

stableEvaluateNodes : ∀ {inputCount stateCount localCount}
  → StableNodes inputCount stateCount localCount
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
stableEvaluateNodes stableNoNodes external state = []
stableEvaluateNodes (nodes ▹ node) external state =
  let previous = stableEvaluateNodes nodes external state
  in snoc previous (stableEvaluateNode external state previous node)

record StableNetlist
  (inputCount outputCount stateCount localCount : ℕ) : Type₀ where
  constructor stableNetlist
  field
    stableInitial : Vec Bit stateCount
    stableNodeSpine : StableNodes inputCount stateCount localCount
    stableOutputs :
      Vec (StableWire inputCount stateCount localCount) outputCount
    stableNext :
      Vec (StableWire inputCount stateCount localCount) stateCount

open StableNetlist public

stableDirectOutputs :
  ∀ {inputCount outputCount stateCount localCount}
  → StableNetlist inputCount outputCount stateCount localCount
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit outputCount
stableDirectOutputs netlist external state =
  stableEvaluateWires external state locals (stableOutputs netlist)
  where
  locals : Vec Bit _
  locals = stableEvaluateNodes (stableNodeSpine netlist) external state

stableDirectNext :
  ∀ {inputCount outputCount stateCount localCount}
  → StableNetlist inputCount outputCount stateCount localCount
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit stateCount
stableDirectNext netlist external state =
  stableEvaluateWires external state locals (stableNext netlist)
  where
  locals : Vec Bit _
  locals = stableEvaluateNodes (stableNodeSpine netlist) external state

-- Compatibility conversion reverses only the meaning of a local index.  Input,
-- state, and literal references are preserved definitionally.

reverseFin : ∀ {count} → Fin count → Fin count
reverseFin {zero} ()
reverseFin {suc count} fzero = fromℕ count
reverseFin {suc count} (fsuc index) = weakenFin (reverseFin index)

convertWire : ∀ {inputCount stateCount localCount}
  → Checked.Wire inputCount stateCount localCount
  → StableWire inputCount stateCount localCount
convertWire (Checked.externalWire index) = stableExternal index
convertWire (Checked.storedWire index) = stableStored index
convertWire (Checked.localWire index) = stableLocal (reverseFin index)
convertWire (Checked.literalWire bit) = stableLiteral bit

convertWires : ∀ {inputCount stateCount localCount count}
  → Vec (Checked.Wire inputCount stateCount localCount) count
  → Vec (StableWire inputCount stateCount localCount) count
convertWires [] = []
convertWires (wire ∷ wires) = convertWire wire ∷ convertWires wires

convertNode : ∀ {inputCount stateCount localCount}
  → Checked.Node inputCount stateCount localCount
  → StableNode inputCount stateCount localCount
convertNode (Checked.invertNode wire) = stableInvert (convertWire wire)
convertNode (Checked.andNode left right) =
  stableAnd (convertWire left) (convertWire right)
convertNode (Checked.orNode left right) =
  stableOr (convertWire left) (convertWire right)
convertNode (Checked.xorNode left right) =
  stableXor (convertWire left) (convertWire right)
convertNode (Checked.muxNode select when-low when-high) =
  stableMux
    (convertWire select) (convertWire when-low) (convertWire when-high)
convertNode (Checked.lutNode table arguments) =
  stableLUT table (convertWires arguments)

convertNodes : ∀ {inputCount stateCount localCount}
  → Checked.Nodes inputCount stateCount localCount
  → StableNodes inputCount stateCount localCount
convertNodes Checked.noNodes = stableNoNodes
convertNodes (nodes Checked.▻ node) = convertNodes nodes ▹ convertNode node

convertNetlist : ∀ {inputCount outputCount stateCount localCount}
  → Checked.CheckedNetlist
      inputCount outputCount stateCount localCount
  → StableNetlist inputCount outputCount stateCount localCount
convertNetlist netlist =
  stableNetlist
    (Checked.checkedInitial netlist)
    (convertNodes (Checked.checkedNodes netlist))
    (convertWires (Checked.checkedOutputs netlist))
    (convertWires (Checked.checkedNext netlist))

lookup-weaken-snoc : ∀ {ℓ} {A : Type ℓ} {count}
  (index : Fin count) (values : Vec A count) (value : A)
  → lookup (weakenFin index) (snoc values value)
    ≡ lookup index values
lookup-weaken-snoc {count = zero} () [] value
lookup-weaken-snoc {count = suc count} fzero (item ∷ items) value = refl
lookup-weaken-snoc {count = suc count} (fsuc index)
  (item ∷ items) value =
  lookup-weaken-snoc index items value

lookup-last-snoc : ∀ {ℓ} {A : Type ℓ} {count}
  (values : Vec A count) (value : A)
  → lookup (fromℕ count) (snoc values value) ≡ value
lookup-last-snoc [] value = refl
lookup-last-snoc (item ∷ items) value = lookup-last-snoc items value

lookup-reverse : ∀ {ℓ} {A : Type ℓ} {count}
  (index : Fin count) (values : Vec A count)
  → lookup (reverseFin index) (reverseValues values)
    ≡ lookup index values
lookup-reverse {count = zero} () []
lookup-reverse {count = suc count} fzero (item ∷ items) =
  lookup-last-snoc (reverseValues items) item
lookup-reverse {count = suc count} (fsuc index) (item ∷ items) =
  lookup-weaken-snoc
    (reverseFin index) (reverseValues items) item
  ∙ lookup-reverse index items

stableWire-convert : ∀ {inputCount stateCount localCount}
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (locals : Vec Bit localCount)
  (wire : Checked.Wire inputCount stateCount localCount)
  → stableEvaluateWire external state
      (reverseValues locals) (convertWire wire)
    ≡ Direct.evaluateWire external state locals wire
stableWire-convert external state locals (Checked.externalWire index) = refl
stableWire-convert external state locals (Checked.storedWire index) = refl
stableWire-convert external state locals (Checked.localWire index) =
  lookup-reverse index locals
stableWire-convert external state locals (Checked.literalWire bit) = refl

stableWires-convert :
  ∀ {inputCount stateCount localCount count}
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (locals : Vec Bit localCount)
  (wires : Vec (Checked.Wire inputCount stateCount localCount) count)
  → stableEvaluateWires external state
      (reverseValues locals) (convertWires wires)
    ≡ Direct.evaluateWires external state locals wires
stableWires-convert external state locals [] = refl
stableWires-convert external state locals (wire ∷ wires) =
  cong₂ _∷_
    (stableWire-convert external state locals wire)
    (stableWires-convert external state locals wires)

stableNode-convert : ∀ {inputCount stateCount localCount}
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (locals : Vec Bit localCount)
  (node : Checked.Node inputCount stateCount localCount)
  → stableEvaluateNode external state
      (reverseValues locals) (convertNode node)
    ≡ Direct.evaluateNode external state locals node
stableNode-convert external state locals (Checked.invertNode wire) =
  cong not (stableWire-convert external state locals wire)
stableNode-convert external state locals (Checked.andNode left right) =
  cong₂ _and_
    (stableWire-convert external state locals left)
    (stableWire-convert external state locals right)
stableNode-convert external state locals (Checked.orNode left right) =
  cong₂ _or_
    (stableWire-convert external state locals left)
    (stableWire-convert external state locals right)
stableNode-convert external state locals (Checked.xorNode left right) =
  cong₂ _⊕_
    (stableWire-convert external state locals left)
    (stableWire-convert external state locals right)
stableNode-convert external state locals
  (Checked.muxNode select when-low when-high) =
  cong₃ mux
    (stableWire-convert external state locals select)
    (stableWire-convert external state locals when-low)
    (stableWire-convert external state locals when-high)
stableNode-convert external state locals (Checked.lutNode table arguments) =
  cong (LUT.evalLUT table)
    (stableWires-convert external state locals arguments)

convertNodes-evaluation : ∀ {inputCount stateCount localCount}
  (nodes : Checked.Nodes inputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → stableEvaluateNodes (convertNodes nodes) external state
    ≡ reverseValues (Direct.evaluateNodes nodes external state)
convertNodes-evaluation Checked.noNodes external state = refl
convertNodes-evaluation (nodes Checked.▻ node) external state =
  let old-agrees = convertNodes-evaluation nodes external state
      direct-old = Direct.evaluateNodes nodes external state
  in cong₂ snoc
      old-agrees
      (cong
        (λ locals →
          stableEvaluateNode external state locals (convertNode node))
        old-agrees
       ∙ stableNode-convert external state direct-old node)

convertWires-evaluation :
  ∀ {inputCount stateCount localCount count}
  (nodes : Checked.Nodes inputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (wires : Vec (Checked.Wire inputCount stateCount localCount) count)
  → stableEvaluateWires external state
      (stableEvaluateNodes (convertNodes nodes) external state)
      (convertWires wires)
    ≡ Direct.evaluateWires external state
        (Direct.evaluateNodes nodes external state) wires
convertWires-evaluation nodes external state wires =
  cong
    (λ locals →
      stableEvaluateWires external state locals (convertWires wires))
    (convertNodes-evaluation nodes external state)
  ∙ stableWires-convert
      external state (Direct.evaluateNodes nodes external state) wires

convertOutputs-evaluation :
  ∀ {inputCount outputCount stateCount localCount}
  (netlist : Checked.CheckedNetlist
    inputCount outputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → stableDirectOutputs (convertNetlist netlist) external state
    ≡ Direct.directOutputs netlist external state
convertOutputs-evaluation netlist external state =
  convertWires-evaluation
    (Checked.checkedNodes netlist) external state
    (Checked.checkedOutputs netlist)

convertNext-evaluation :
  ∀ {inputCount outputCount stateCount localCount}
  (netlist : Checked.CheckedNetlist
    inputCount outputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → stableDirectNext (convertNetlist netlist) external state
    ≡ Direct.directNext netlist external state
convertNext-evaluation netlist external state =
  convertWires-evaluation
    (Checked.checkedNodes netlist) external state
    (Checked.checkedNext netlist)

convertOutputs-compile :
  ∀ {inputCount outputCount stateCount localCount}
  (netlist : Checked.CheckedNetlist
    inputCount outputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → stableDirectOutputs (convertNetlist netlist) external state
    ≡ Semantics.observe (Checked.compileNetlist netlist) external state
convertOutputs-compile netlist external state =
  convertOutputs-evaluation netlist external state
  ∙ Direct.directOutputs-compile netlist external state

convertNext-compile :
  ∀ {inputCount outputCount stateCount localCount}
  (netlist : Checked.CheckedNetlist
    inputCount outputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → stableDirectNext (convertNetlist netlist) external state
    ≡ Semantics.step
        (Checked.compileNetlist netlist)
        Semantics.risingEdge external state
convertNext-compile netlist external state =
  convertNext-evaluation netlist external state
  ∙ Direct.directNext-compile netlist external state

stableAppend-preserves-local : ∀ {inputCount stateCount localCount}
  (nodes : StableNodes inputCount stateCount localCount)
  (node : StableNode inputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (index : Fin localCount)
  → stableEvaluateWire external state
      (stableEvaluateNodes (nodes ▹ node) external state)
      (stableLocal (weakenFin index))
    ≡ stableEvaluateWire external state
        (stableEvaluateNodes nodes external state)
        (stableLocal index)
stableAppend-preserves-local nodes node external state index =
  lookup-weaken-snoc index
    (stableEvaluateNodes nodes external state)
    (stableEvaluateNode external state
      (stableEvaluateNodes nodes external state) node)

stableAppend-new-index : ∀ {inputCount stateCount localCount}
  (nodes : StableNodes inputCount stateCount localCount)
  (node : StableNode inputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → stableEvaluateWire external state
      (stableEvaluateNodes (nodes ▹ node) external state)
      (stableLocal (fromℕ localCount))
    ≡ stableEvaluateNode external state
        (stableEvaluateNodes nodes external state) node
stableAppend-new-index nodes node external state =
  lookup-last-snoc
    (stableEvaluateNodes nodes external state)
    (stableEvaluateNode external state
      (stableEvaluateNodes nodes external state) node)