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

module Spartan6.Netlist.DAGEvaluation where

open import Spartan6.Prelude

import Spartan6.Netlist.BuilderCore as Builder
import Spartan6.Netlist.Checked as Checked
import Spartan6.Netlist.Expression as Expression
import Spartan6.Netlist.Raw as Raw
import Spartan6.Primitive.LUT as LUT
import Spartan6.Semantics.Design as Semantics

-- Direct evaluation retains the node spine as a shared environment.  The
-- newest checked node is consed at local index zero, exactly matching the
-- existing newest-at-zero `Checked.Nodes` representation.

evaluateWire : ∀ {inputCount stateCount localCount}
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
  → Checked.Wire inputCount stateCount localCount
  → Bit
evaluateWire external state locals (Checked.externalWire index) =
  lookup index external
evaluateWire external state locals (Checked.storedWire index) =
  lookup index state
evaluateWire external state locals (Checked.localWire index) =
  lookup index locals
evaluateWire external state locals (Checked.literalWire bit) = bit

evaluateWires : ∀ {inputCount stateCount localCount count}
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
  → Vec (Checked.Wire inputCount stateCount localCount) count
  → Vec Bit count
evaluateWires external state locals [] = []
evaluateWires external state locals (wire ∷ wires) =
  evaluateWire external state locals wire
  ∷ evaluateWires external state locals wires

evaluateNode : ∀ {inputCount stateCount localCount}
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
  → Checked.Node inputCount stateCount localCount
  → Bit
evaluateNode external state locals (Checked.invertNode wire) =
  not (evaluateWire external state locals wire)
evaluateNode external state locals (Checked.andNode left right) =
  evaluateWire external state locals left
  and evaluateWire external state locals right
evaluateNode external state locals (Checked.orNode left right) =
  evaluateWire external state locals left
  or evaluateWire external state locals right
evaluateNode external state locals (Checked.xorNode left right) =
  evaluateWire external state locals left
  ⊕ evaluateWire external state locals right
evaluateNode external state locals
  (Checked.muxNode select when-low when-high) =
  mux (evaluateWire external state locals select)
      (evaluateWire external state locals when-low)
      (evaluateWire external state locals when-high)
evaluateNode external state locals (Checked.lutNode table arguments) =
  LUT.evalLUT table (evaluateWires external state locals arguments)

evaluateNodes : ∀ {inputCount stateCount localCount}
  → Checked.Nodes inputCount stateCount localCount
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit localCount
evaluateNodes Checked.noNodes external state = []
evaluateNodes (nodes Checked.▻ node) external state =
  let previous = evaluateNodes nodes external state
  in evaluateNode external state previous node ∷ previous

directOutputs : ∀ {inputCount outputCount stateCount localCount}
  → Checked.CheckedNetlist
      inputCount outputCount stateCount localCount
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit outputCount
directOutputs netlist external state =
  evaluateWires external state locals (Checked.checkedOutputs netlist)
  where
  locals : Vec Bit _
  locals = evaluateNodes (Checked.checkedNodes netlist) external state

directNext : ∀ {inputCount outputCount stateCount localCount}
  → Checked.CheckedNetlist
      inputCount outputCount stateCount localCount
  → Vec Bit inputCount
  → Vec Bit stateCount
  → Vec Bit stateCount
directNext netlist external state =
  evaluateWires external state locals (Checked.checkedNext netlist)
  where
  locals : Vec Bit _
  locals = evaluateNodes (Checked.checkedNodes netlist) external state

-- The proof is parameterized by an arbitrary agreement between a direct local
-- environment and compiled local expressions.  This avoids a circular proof
-- between local-wire lookup and node-spine evaluation.

lookup-evalAll : ∀ {inputCount stateCount count}
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (expressions : Vec (Expression.Expr inputCount stateCount) count)
  (index : Fin count)
  → lookup index (Expression.evalAll external state expressions)
    ≡ Expression.eval external state (lookup index expressions)
lookup-evalAll external state (expression ∷ expressions) fzero = refl
lookup-evalAll external state (expression ∷ expressions) (fsuc index) =
  lookup-evalAll external state expressions index

evaluateWire-compiled :
  ∀ {inputCount stateCount localCount}
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (locals : Vec Bit localCount)
  (expressions : Vec (Expression.Expr inputCount stateCount) localCount)
  → locals ≡ Expression.evalAll external state expressions
  → (wire : Checked.Wire inputCount stateCount localCount)
  → evaluateWire external state locals wire
    ≡ Expression.eval external state
        (Checked.compileWire expressions wire)
evaluateWire-compiled external state locals expressions agreement
  (Checked.externalWire index) = refl
evaluateWire-compiled external state locals expressions agreement
  (Checked.storedWire index) = refl
evaluateWire-compiled external state locals expressions agreement
  (Checked.localWire index) =
  cong (lookup index) agreement
  ∙ lookup-evalAll external state expressions index
evaluateWire-compiled external state locals expressions agreement
  (Checked.literalWire bit) = refl

evaluateWires-compiled :
  ∀ {inputCount stateCount localCount count}
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (locals : Vec Bit localCount)
  (expressions : Vec (Expression.Expr inputCount stateCount) localCount)
  → locals ≡ Expression.evalAll external state expressions
  → (wires :
      Vec (Checked.Wire inputCount stateCount localCount) count)
  → evaluateWires external state locals wires
    ≡ Expression.evalAll external state
        (Checked.compileWires expressions wires)
evaluateWires-compiled external state locals expressions agreement [] = refl
evaluateWires-compiled external state locals expressions agreement
  (wire ∷ wires) =
  cong₂ _∷_
    (evaluateWire-compiled
      external state locals expressions agreement wire)
    (evaluateWires-compiled
      external state locals expressions agreement wires)

evaluateNode-compiled :
  ∀ {inputCount stateCount localCount}
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (locals : Vec Bit localCount)
  (expressions : Vec (Expression.Expr inputCount stateCount) localCount)
  → locals ≡ Expression.evalAll external state expressions
  → (node : Checked.Node inputCount stateCount localCount)
  → evaluateNode external state locals node
    ≡ Expression.eval external state
        (Checked.compileNode expressions node)
evaluateNode-compiled external state locals expressions agreement
  (Checked.invertNode wire) =
  cong not
    (evaluateWire-compiled
      external state locals expressions agreement wire)
evaluateNode-compiled external state locals expressions agreement
  (Checked.andNode left right) =
  cong₂ _and_
    (evaluateWire-compiled
      external state locals expressions agreement left)
    (evaluateWire-compiled
      external state locals expressions agreement right)
evaluateNode-compiled external state locals expressions agreement
  (Checked.orNode left right) =
  cong₂ _or_
    (evaluateWire-compiled
      external state locals expressions agreement left)
    (evaluateWire-compiled
      external state locals expressions agreement right)
evaluateNode-compiled external state locals expressions agreement
  (Checked.xorNode left right) =
  cong₂ _⊕_
    (evaluateWire-compiled
      external state locals expressions agreement left)
    (evaluateWire-compiled
      external state locals expressions agreement right)
evaluateNode-compiled external state locals expressions agreement
  (Checked.muxNode select when-low when-high) =
  cong₃ mux
    (evaluateWire-compiled
      external state locals expressions agreement select)
    (evaluateWire-compiled
      external state locals expressions agreement when-low)
    (evaluateWire-compiled
      external state locals expressions agreement when-high)
evaluateNode-compiled external state locals expressions agreement
  (Checked.lutNode table arguments) =
  cong (LUT.evalLUT table)
    (evaluateWires-compiled
      external state locals expressions agreement arguments)

evaluateNodes-compiled : ∀ {inputCount stateCount localCount}
  (nodes : Checked.Nodes inputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → evaluateNodes nodes external state
    ≡ Expression.evalAll external state (Checked.compileNodes nodes)
evaluateNodes-compiled Checked.noNodes external state = refl
evaluateNodes-compiled (nodes Checked.▻ node) external state =
  let previous-agrees = evaluateNodes-compiled nodes external state
  in cong₂ _∷_
      (evaluateNode-compiled
        external state
        (evaluateNodes nodes external state)
        (Checked.compileNodes nodes)
        previous-agrees node)
      previous-agrees

directOutputs-compile :
  ∀ {inputCount outputCount stateCount localCount}
  (netlist : Checked.CheckedNetlist
    inputCount outputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → directOutputs netlist external state
    ≡ Semantics.observe (Checked.compileNetlist netlist) external state
directOutputs-compile netlist external state =
  evaluateWires-compiled
    external state
    (evaluateNodes (Checked.checkedNodes netlist) external state)
    (Checked.compileNodes (Checked.checkedNodes netlist))
    (evaluateNodes-compiled
      (Checked.checkedNodes netlist) external state)
    (Checked.checkedOutputs netlist)

directNext-compile :
  ∀ {inputCount outputCount stateCount localCount}
  (netlist : Checked.CheckedNetlist
    inputCount outputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → directNext netlist external state
    ≡ Semantics.step
        (Checked.compileNetlist netlist)
        Semantics.risingEdge external state
directNext-compile netlist external state =
  evaluateWires-compiled
    external state
    (evaluateNodes (Checked.checkedNodes netlist) external state)
    (Checked.compileNodes (Checked.checkedNodes netlist))
    (evaluateNodes-compiled
      (Checked.checkedNodes netlist) external state)
    (Checked.checkedNext netlist)

-- Extension probes.  The first theorem is representation-local; the second is
-- the corresponding existing BuilderCore operation.  Both are constructor
-- proofs, so an ordinary extension never replays the old node spine.

append-preserves-wire : ∀ {inputCount stateCount localCount}
  (nodes : Checked.Nodes inputCount stateCount localCount)
  (node : Checked.Node inputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (wire : Checked.Wire inputCount stateCount localCount)
  → evaluateWire external state
      (evaluateNodes (nodes Checked.▻ node) external state)
      (Builder.liftLocalWire wire)
    ≡ evaluateWire external state
        (evaluateNodes nodes external state) wire
append-preserves-wire nodes node external state
  (Checked.externalWire index) = refl
append-preserves-wire nodes node external state
  (Checked.storedWire index) = refl
append-preserves-wire nodes node external state
  (Checked.localWire index) = refl
append-preserves-wire nodes node external state
  (Checked.literalWire bit) = refl

append-newest-value : ∀ {inputCount stateCount localCount}
  (nodes : Checked.Nodes inputCount stateCount localCount)
  (node : Checked.Node inputCount stateCount localCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → evaluateWire external state
      (evaluateNodes (nodes Checked.▻ node) external state)
      (Checked.localWire fzero)
    ≡ evaluateNode external state
        (evaluateNodes nodes external state) node
append-newest-value nodes node external state = refl

evaluateBuilderWire : ∀ {inputCount stateCount}
  (current : Builder.Builder inputCount stateCount)
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  → Checked.Wire inputCount stateCount (Builder.builderLocalCount current)
  → Bit
evaluateBuilderWire current external state =
  evaluateWire external state
    (evaluateNodes (Builder.builderNodes current) external state)

extendBuilder-preserves-wire : ∀ {inputCount stateCount}
  (current : Builder.Builder inputCount stateCount)
  (net-id : Raw.NetId)
  (node : Checked.Node inputCount stateCount
    (Builder.builderLocalCount current))
  (external : Vec Bit inputCount)
  (state : Vec Bit stateCount)
  (wire : Checked.Wire inputCount stateCount
    (Builder.builderLocalCount current))
  → evaluateBuilderWire
      (Builder.extendBuilder current net-id node)
      external state (Builder.liftLocalWire wire)
    ≡ evaluateBuilderWire current external state wire
extendBuilder-preserves-wire
  (Builder.builder localCount nodes bindings)
  net-id node external state (Checked.externalWire index) = refl
extendBuilder-preserves-wire
  (Builder.builder localCount nodes bindings)
  net-id node external state (Checked.storedWire index) = refl
extendBuilder-preserves-wire
  (Builder.builder localCount nodes bindings)
  net-id node external state (Checked.localWire index) = refl
extendBuilder-preserves-wire
  (Builder.builder localCount nodes bindings)
  net-id node external state (Checked.literalWire bit) = refl