{-# LANGUAGE ScopedTypeVariables #-}

-- | A `UplcEvaluator` shared between the regular and steppable CEK machines.
module PlutusConformance.CekEvaluator (mkCekEvaluator) where

import PlutusConformance.Common
  ( EvaluationResult (..)
  , UplcEvaluator (..)
  )
import PlutusCore.Default
  ( DefaultFun
  , DefaultUni
  )
import PlutusCore.Evaluation.Machine.MachineParameters qualified as UPLC
import PlutusCore.Evaluation.Machine.MachineParameters.Default
  ( mkMachineVariantParametersFor
  )
import PlutusCore.Name.Unique (Name)
import PlutusPrelude (def)
import UntypedPlutusCore qualified as UPLC
import UntypedPlutusCore.Evaluation.Machine.Cek
  ( CekEvaluationException
  , CekMachineCosts
  , CekValue
  , CountingSt (..)
  , ExBudgetMode
  , counting
  )

{-| Build a `UplcEvaluator` around a CEK-machine "run" function shaped like
`UntypedPlutusCore.Evaluation.Machine.Cek.runCekNoEmit`.  Both the regular CEK
machine and the steppable one
(`UntypedPlutusCore.Evaluation.Machine.SteppableCek.runCekNoEmit`) have exactly
this type (the latter's haddock says it "provides the same interface to the
original CEK machine") so this is shared between the `haskell-conformance` and
`haskell-steppable-conformance` test suites rather than being duplicated. -}
mkCekEvaluator
  :: ( UPLC.MachineParameters CekMachineCosts DefaultFun (CekValue DefaultUni DefaultFun ())
       -> ExBudgetMode CountingSt DefaultUni DefaultFun
       -> UPLC.Term Name DefaultUni DefaultFun ()
       -> ( Either (CekEvaluationException Name DefaultUni DefaultFun) (UPLC.Term Name DefaultUni DefaultFun ())
          , CountingSt
          )
     )
  -> UplcEvaluator
mkCekEvaluator :: (MachineParameters
   CekMachineCosts DefaultFun (CekValue DefaultUni DefaultFun ())
 -> ExBudgetMode CountingSt DefaultUni DefaultFun
 -> Term Name DefaultUni DefaultFun ()
 -> (Either
       (CekEvaluationException Name DefaultUni DefaultFun)
       (Term Name DefaultUni DefaultFun ()),
     CountingSt))
-> UplcEvaluator
mkCekEvaluator MachineParameters
  CekMachineCosts DefaultFun (CekValue DefaultUni DefaultFun ())
-> ExBudgetMode CountingSt DefaultUni DefaultFun
-> Term Name DefaultUni DefaultFun ()
-> (Either
      (CekEvaluationException Name DefaultUni DefaultFun)
      (Term Name DefaultUni DefaultFun ()),
    CountingSt)
runCekNoEmit = (CostModelParams -> UplcEvaluatorFun (UplcProg, ExBudget))
-> UplcEvaluator
UplcEvaluatorWithCosting ((CostModelParams -> UplcEvaluatorFun (UplcProg, ExBudget))
 -> UplcEvaluator)
-> (CostModelParams -> UplcEvaluatorFun (UplcProg, ExBudget))
-> UplcEvaluator
forall a b. (a -> b) -> a -> b
$ \CostModelParams
modelParams (UPLC.Program ()
a Version
v Term Name DefaultUni DefaultFun ()
t) ->
  case [BuiltinSemanticsVariant DefaultFun]
-> CostModelParams
-> Either
     CostModelApplyError
     [(BuiltinSemanticsVariant DefaultFun,
       DefaultMachineVariantParameters)]
forall (m :: * -> *).
MonadError CostModelApplyError m =>
[BuiltinSemanticsVariant DefaultFun]
-> CostModelParams
-> m [(BuiltinSemanticsVariant DefaultFun,
       DefaultMachineVariantParameters)]
mkMachineVariantParametersFor [BuiltinSemanticsVariant DefaultFun
forall a. Default a => a
def] CostModelParams
modelParams of
    Left CostModelApplyError
_ -> EvaluationResult (UplcProg, ExBudget)
forall res. EvaluationResult res
BadMachineParameters
    Right [(BuiltinSemanticsVariant DefaultFun,
  DefaultMachineVariantParameters)]
machParamsList ->
      case BuiltinSemanticsVariant DefaultFun
-> [(BuiltinSemanticsVariant DefaultFun,
     DefaultMachineVariantParameters)]
-> Maybe DefaultMachineVariantParameters
forall a b. Eq a => a -> [(a, b)] -> Maybe b
lookup BuiltinSemanticsVariant DefaultFun
forall a. Default a => a
def [(BuiltinSemanticsVariant DefaultFun,
  DefaultMachineVariantParameters)]
machParamsList of
        Maybe DefaultMachineVariantParameters
Nothing -> EvaluationResult (UplcProg, ExBudget)
forall res. EvaluationResult res
BadMachineParameters
        Just DefaultMachineVariantParameters
p ->
          let params :: MachineParameters
  CekMachineCosts DefaultFun (CekValue DefaultUni DefaultFun ())
params = CaserBuiltin (UniOf (CekValue DefaultUni DefaultFun ()))
-> DefaultMachineVariantParameters
-> MachineParameters
     CekMachineCosts DefaultFun (CekValue DefaultUni DefaultFun ())
forall machineCosts fun val.
CaserBuiltin (UniOf val)
-> MachineVariantParameters machineCosts fun val
-> MachineParameters machineCosts fun val
UPLC.MachineParameters CaserBuiltin (UniOf (CekValue DefaultUni DefaultFun ()))
CaserBuiltin DefaultUni
forall a. Default a => a
def DefaultMachineVariantParameters
p
           in -- runCek-like functions (e.g. evaluateCekNoEmit) are partial on term's with
              -- free variables, that is why we manually check first for any free vars
              case Term Name DefaultUni DefaultFun ()
-> Either
     FreeVariableError (Term NamedDeBruijn DefaultUni DefaultFun ())
forall (m :: * -> *) (uni :: * -> *) fun ann.
MonadError FreeVariableError m =>
Term Name uni fun ann -> m (Term NamedDeBruijn uni fun ann)
UPLC.deBruijnTerm Term Name DefaultUni DefaultFun ()
t of
                Left (FreeVariableError
_ :: UPLC.FreeVariableError) -> EvaluationResult (UplcProg, ExBudget)
forall res. EvaluationResult res
DecodeError -- For consistency with the flat decoder.
                Right Term NamedDeBruijn DefaultUni DefaultFun ()
_ ->
                  case MachineParameters
  CekMachineCosts DefaultFun (CekValue DefaultUni DefaultFun ())
-> ExBudgetMode CountingSt DefaultUni DefaultFun
-> Term Name DefaultUni DefaultFun ()
-> (Either
      (CekEvaluationException Name DefaultUni DefaultFun)
      (Term Name DefaultUni DefaultFun ()),
    CountingSt)
runCekNoEmit MachineParameters
  CekMachineCosts DefaultFun (CekValue DefaultUni DefaultFun ())
params ExBudgetMode CountingSt DefaultUni DefaultFun
forall (uni :: * -> *) fun. ExBudgetMode CountingSt uni fun
counting Term Name DefaultUni DefaultFun ()
t of
                    (Left CekEvaluationException Name DefaultUni DefaultFun
_, CountingSt
_) -> EvaluationResult (UplcProg, ExBudget)
forall res. EvaluationResult res
EvalFailure
                    (Right Term Name DefaultUni DefaultFun ()
prog, CountingSt ExBudget
cost) -> (UplcProg, ExBudget) -> EvaluationResult (UplcProg, ExBudget)
forall res. res -> EvaluationResult res
EvalSuccess (() -> Version -> Term Name DefaultUni DefaultFun () -> UplcProg
forall name (uni :: * -> *) fun ann.
ann -> Version -> Term name uni fun ann -> Program name uni fun ann
UPLC.Program ()
a Version
v Term Name DefaultUni DefaultFun ()
prog, ExBudget
cost)