{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeOperators #-}

{-| A top-down pass converting 2 or more consecutive casing on lists,
if the heads are all unused, and the tails are all unused except for
being immediately matched on.

Example:

case xs of _h1 t1 ->
  case t1 of _h2 t2 ->
    case t2 of _h3 t3 ->
      case t3 of _h4 t4 ->
        case t4 of h5 t5 -> ...t5...

===>

case (drop 4 xs) of h5 t5 -> ...t5... -}
module PlutusIR.Transform.CollapseCase
  ( collapseCase
  , collapseCasePassSC
  ) where

import PlutusCore qualified as PLC
import PlutusCore.Analysis.Usages qualified as Usages
import PlutusCore.Annotation
import PlutusCore.Name.Unique
import PlutusIR
import PlutusIR.Pass
import PlutusIR.Subst (termUsages)
import PlutusIR.Transform.Rename ()
import PlutusIR.TypeCheck qualified as TC

import Control.Lens (over, transformOf, view)

collapseCasePassSC
  :: (uni ~ PLC.DefaultUni, fun ~ PLC.DefaultFun, Applicative m, AnnCase a)
  => TC.PirTCConfig uni fun
  -> Pass m TyName Name uni fun a
collapseCasePassSC :: forall (uni :: * -> *) fun (m :: * -> *) a.
(uni ~ DefaultUni, fun ~ DefaultFun, Applicative m, AnnCase a) =>
PirTCConfig uni fun -> Pass m TyName Name uni fun a
collapseCasePassSC PirTCConfig uni fun
tcconfig =
  String
-> Pass m TyName Name uni fun a -> Pass m TyName Name uni fun a
forall (m :: * -> *) tyname name (uni :: * -> *) fun a.
String
-> Pass m tyname name uni fun a -> Pass m tyname name uni fun a
NamedPass String
"collapse cases on lists into dropLists" (Pass m TyName Name uni fun a -> Pass m TyName Name uni fun a)
-> Pass m TyName Name uni fun a -> Pass m TyName Name uni fun a
forall a b. (a -> b) -> a -> b
$
    (Term TyName Name uni fun a -> m (Term TyName Name uni fun a))
-> [Condition TyName Name uni fun a]
-> [BiCondition TyName Name uni fun a]
-> Pass m TyName Name uni fun a
forall (m :: * -> *) tyname name (uni :: * -> *) fun a.
(Term tyname name uni fun a -> m (Term tyname name uni fun a))
-> [Condition tyname name uni fun a]
-> [BiCondition tyname name uni fun a]
-> Pass m tyname name uni fun a
Pass
      (Term TyName Name uni fun a -> m (Term TyName Name uni fun a)
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Term TyName Name uni fun a -> m (Term TyName Name uni fun a))
-> (Term TyName Name uni fun a -> Term TyName Name uni fun a)
-> Term TyName Name uni fun a
-> m (Term TyName Name uni fun a)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Term TyName Name uni fun a -> Term TyName Name uni fun a
forall name (uni :: * -> *) fun a.
(uni ~ DefaultUni, fun ~ DefaultFun, HasUnique name TermUnique,
 AnnCase a) =>
Term TyName name uni fun a -> Term TyName name uni fun a
collapseCase)
      [PirTCConfig uni fun -> Condition TyName Name uni fun a
forall (uni :: * -> *) fun a.
(Typecheckable uni fun, GEq uni) =>
PirTCConfig uni fun -> Condition TyName Name uni fun a
Typechecks PirTCConfig uni fun
tcconfig]
      [Condition TyName Name uni fun a
-> BiCondition TyName Name uni fun a
forall tyname name (uni :: * -> *) fun a.
Condition tyname name uni fun a
-> BiCondition tyname name uni fun a
ConstCondition (PirTCConfig uni fun -> Condition TyName Name uni fun a
forall (uni :: * -> *) fun a.
(Typecheckable uni fun, GEq uni) =>
PirTCConfig uni fun -> Condition TyName Name uni fun a
Typechecks PirTCConfig uni fun
tcconfig)]

collapseCase
  :: forall name uni fun a
   . ( uni ~ PLC.DefaultUni
     , fun ~ PLC.DefaultFun
     , HasUnique name TermUnique
     , AnnCase a
     )
  => Term TyName name uni fun a
  -> Term TyName name uni fun a
collapseCase :: forall name (uni :: * -> *) fun a.
(uni ~ DefaultUni, fun ~ DefaultFun, HasUnique name TermUnique,
 AnnCase a) =>
Term TyName name uni fun a -> Term TyName name uni fun a
collapseCase Term TyName name uni fun a
t = Term TyName name uni fun a
-> (Term TyName name uni fun a -> Term TyName name uni fun a)
-> Maybe (Term TyName name uni fun a)
-> Term TyName name uni fun a
forall b a. b -> (a -> b) -> Maybe a -> b
maybe (ASetter
  (Term TyName name uni fun a)
  (Term TyName name uni fun a)
  (Term TyName name uni fun a)
  (Term TyName name uni fun a)
-> (Term TyName name uni fun a -> Term TyName name uni fun a)
-> Term TyName name uni fun a
-> Term TyName name uni fun a
forall s t a b. ASetter s t a b -> (a -> b) -> s -> t
over ASetter
  (Term TyName name uni fun a)
  (Term TyName name uni fun a)
  (Term TyName name uni fun a)
  (Term TyName name uni fun a)
forall tyname name (uni :: * -> *) fun a (f :: * -> *).
Applicative f =>
(Term tyname name uni fun a -> f (Term tyname name uni fun a))
-> Term tyname name uni fun a -> f (Term tyname name uni fun a)
termSubterms Term TyName name uni fun a -> Term TyName name uni fun a
forall name (uni :: * -> *) fun a.
(uni ~ DefaultUni, fun ~ DefaultFun, HasUnique name TermUnique,
 AnnCase a) =>
Term TyName name uni fun a -> Term TyName name uni fun a
collapseCase Term TyName name uni fun a
t) Term TyName name uni fun a -> Term TyName name uni fun a
forall name (uni :: * -> *) fun a.
(uni ~ DefaultUni, fun ~ DefaultFun, HasUnique name TermUnique,
 AnnCase a) =>
Term TyName name uni fun a -> Term TyName name uni fun a
collapseCase (Term TyName name uni fun a -> Maybe (Term TyName name uni fun a)
collapse Term TyName name uni fun a
t)
  where
    collapse :: Term TyName name uni fun a -> Maybe (Term TyName name uni fun a)
collapse = \case
      -- First casing in the sequence - go from here
      Case a
a Type TyName uni a
_resTy Term TyName name uni fun a
scrut [LamAbs a
_ name
hd Type TyName uni a
elemTy (LamAbs a
_ name
tl Type TyName uni a
_ Term TyName name uni fun a
body)]
        | a -> Bool
forall a. AnnCase a => a -> Bool
annIsSafeToDrop a
a
        , name -> Usages -> Int
forall n unique. HasUnique n unique => n -> Usages -> Int
Usages.getUsageCount name
hd (Term TyName name uni fun a -> Usages
forall name tyname (uni :: * -> *) fun ann.
(HasUnique name TermUnique, HasUnique tyname TypeUnique) =>
Term tyname name uni fun ann -> Usages
termUsages Term TyName name uni fun a
body) Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 ->
            Integer
-> a
-> Term TyName name uni fun a
-> Type TyName uni a
-> name
-> Term TyName name uni fun a
-> Maybe (Term TyName name uni fun a)
go (Integer
1 :: Integer) (Term TyName name uni fun a -> a
forall a. Term TyName name uni fun a -> a
forall (f :: * -> *) a. HasAnn f => f a -> a
getAnn Term TyName name uni fun a
scrut) Term TyName name uni fun a
scrut Type TyName uni a
elemTy name
tl Term TyName name uni fun a
body
      Term TyName name uni fun a
_ -> Maybe (Term TyName name uni fun a)
forall a. Maybe a
Nothing

    go
      :: Integer
      -> a
      -- \^ annotation on the first scrutinee in the sequence
      -> Term TyName name uni fun a
      -- \^ first scrutinee in the sequence
      -> Type TyName uni a
      -- \^ list element type
      -> name
      -- \^ current tail
      -> Term TyName name uni fun a
      -- \^ current body
      -> Maybe (Term TyName name uni fun a)
    go :: Integer
-> a
-> Term TyName name uni fun a
-> Type TyName uni a
-> name
-> Term TyName name uni fun a
-> Maybe (Term TyName name uni fun a)
go Integer
k a
aTop Term TyName name uni fun a
scrutTop Type TyName uni a
elemTy name
tl0 Term TyName name uni fun a
body0 = case Term TyName name uni fun a
body0 of
      Case a
a Type TyName uni a
_resTy (Var a
_ name
scrut) [LamAbs a
_ name
hd Type TyName uni a
_tyElem' (LamAbs a
_ name
tl Type TyName uni a
_ Term TyName name uni fun a
body)]
        | a -> Bool
forall a. AnnCase a => a -> Bool
annIsSafeToDrop a
a
        , Getting Unique name Unique -> name -> Unique
forall s (m :: * -> *) a. MonadReader s m => Getting a s a -> m a
view Getting Unique name Unique
forall name unique. HasUnique name unique => Lens' name Unique
Lens' name Unique
theUnique name
scrut Unique -> Unique -> Bool
forall a. Eq a => a -> a -> Bool
== Getting Unique name Unique -> name -> Unique
forall s (m :: * -> *) a. MonadReader s m => Getting a s a -> m a
view Getting Unique name Unique
forall name unique. HasUnique name unique => Lens' name Unique
Lens' name Unique
theUnique name
tl0
        , let usages :: Usages
usages = Term TyName name uni fun a -> Usages
forall name tyname (uni :: * -> *) fun ann.
(HasUnique name TermUnique, HasUnique tyname TypeUnique) =>
Term tyname name uni fun ann -> Usages
termUsages Term TyName name uni fun a
body
        , -- new head must be unused in the body for the sequence to continue
          name -> Usages -> Int
forall n unique. HasUnique n unique => n -> Usages -> Int
Usages.getUsageCount name
hd Usages
usages Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0
        , -- original tail must be unused in the body for the sequence to continue
          name -> Usages -> Int
forall n unique. HasUnique n unique => n -> Usages -> Int
Usages.getUsageCount name
tl0 Usages
usages Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 ->
            -- recursive with new tail and new body
            Integer
-> a
-> Term TyName name uni fun a
-> Type TyName uni a
-> name
-> Term TyName name uni fun a
-> Maybe (Term TyName name uni fun a)
go (Integer
k Integer -> Integer -> Integer
forall a. Num a => a -> a -> a
+ Integer
1) a
aTop Term TyName name uni fun a
scrutTop Type TyName uni a
elemTy name
tl Term TyName name uni fun a
body
      Term TyName name uni fun a
_
        | name -> Usages -> Int
forall n unique. HasUnique n unique => n -> Usages -> Int
Usages.getUsageCount name
tl0 (Term TyName name uni fun a -> Usages
forall name tyname (uni :: * -> *) fun ann.
(HasUnique name TermUnique, HasUnique tyname TypeUnique) =>
Term tyname name uni fun ann -> Usages
termUsages Term TyName name uni fun a
body0) Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
1
        , Integer
k Integer -> Integer -> Bool
forall a. Ord a => a -> a -> Bool
>= Integer
2 ->
            Term TyName name uni fun a -> Maybe (Term TyName name uni fun a)
forall a. a -> Maybe a
Just (Term TyName name uni fun a -> Maybe (Term TyName name uni fun a))
-> Term TyName name uni fun a -> Maybe (Term TyName name uni fun a)
forall a b. (a -> b) -> a -> b
$ name
-> Term TyName name uni fun a
-> Term TyName name uni fun a
-> Term TyName name uni fun a
forall {unique} {unique} {s} {s} {tyname} {uni :: * -> *} {fun}
       {a}.
(Coercible unique Int, Coercible unique Int, HasUnique s unique,
 HasUnique s unique) =>
s
-> Term tyname s uni fun a
-> Term tyname s uni fun a
-> Term tyname s uni fun a
substVar name
tl0 Term TyName name uni fun a
dropped Term TyName name uni fun a
body0
        | Bool
otherwise -> Maybe (Term TyName name uni fun a)
forall a. Maybe a
Nothing
      where
        dropped :: Term TyName name uni fun a
dropped = a
-> Integer
-> Type TyName uni a
-> Term TyName name uni fun a
-> Term TyName name uni fun a
mkDrop a
aTop Integer
k Type TyName uni a
elemTy Term TyName name uni fun a
scrutTop

        substVar :: s
-> Term tyname s uni fun a
-> Term tyname s uni fun a
-> Term tyname s uni fun a
substVar s
n Term tyname s uni fun a
new = ASetter
  (Term tyname s uni fun a)
  (Term tyname s uni fun a)
  (Term tyname s uni fun a)
  (Term tyname s uni fun a)
-> (Term tyname s uni fun a -> Term tyname s uni fun a)
-> Term tyname s uni fun a
-> Term tyname s uni fun a
forall a b. ASetter a b a b -> (b -> b) -> a -> b
transformOf ASetter
  (Term tyname s uni fun a)
  (Term tyname s uni fun a)
  (Term tyname s uni fun a)
  (Term tyname s uni fun a)
forall tyname name (uni :: * -> *) fun a (f :: * -> *).
Applicative f =>
(Term tyname name uni fun a -> f (Term tyname name uni fun a))
-> Term tyname name uni fun a -> f (Term tyname name uni fun a)
termSubterms ((Term tyname s uni fun a -> Term tyname s uni fun a)
 -> Term tyname s uni fun a -> Term tyname s uni fun a)
-> (Term tyname s uni fun a -> Term tyname s uni fun a)
-> Term tyname s uni fun a
-> Term tyname s uni fun a
forall a b. (a -> b) -> a -> b
$ \case
          Var a
_ s
v | Getting Unique s Unique -> s -> Unique
forall s (m :: * -> *) a. MonadReader s m => Getting a s a -> m a
view Getting Unique s Unique
forall name unique. HasUnique name unique => Lens' name Unique
Lens' s Unique
theUnique s
v Unique -> Unique -> Bool
forall a. Eq a => a -> a -> Bool
== Getting Unique s Unique -> s -> Unique
forall s (m :: * -> *) a. MonadReader s m => Getting a s a -> m a
view Getting Unique s Unique
forall name unique. HasUnique name unique => Lens' name Unique
Lens' s Unique
theUnique s
n -> Term tyname s uni fun a
new
          Term tyname s uni fun a
other -> Term tyname s uni fun a
other

    mkDrop
      :: a
      -> Integer
      -> Type TyName uni a
      -> Term TyName name uni fun a
      -> Term TyName name uni fun a
    mkDrop :: a
-> Integer
-> Type TyName uni a
-> Term TyName name uni fun a
-> Term TyName name uni fun a
mkDrop a
ann Integer
k Type TyName uni a
elemTy =
      a
-> Term TyName name uni fun a
-> Term TyName name uni fun a
-> Term TyName name uni fun a
forall tyname name (uni :: * -> *) fun a.
a
-> Term tyname name uni fun a
-> Term tyname name uni fun a
-> Term tyname name uni fun a
Apply
        a
ann
        ( a
-> Term TyName name uni fun a
-> Term TyName name uni fun a
-> Term TyName name uni fun a
forall tyname name (uni :: * -> *) fun a.
a
-> Term tyname name uni fun a
-> Term tyname name uni fun a
-> Term tyname name uni fun a
Apply
            a
ann
            (a
-> Term TyName name uni fun a
-> Type TyName uni a
-> Term TyName name uni fun a
forall tyname name (uni :: * -> *) fun a.
a
-> Term tyname name uni fun a
-> Type tyname uni a
-> Term tyname name uni fun a
TyInst a
ann (a -> fun -> Term TyName name uni fun a
forall tyname name (uni :: * -> *) fun a.
a -> fun -> Term tyname name uni fun a
Builtin a
ann fun
DefaultFun
PLC.DropList) Type TyName uni a
elemTy)
            (a -> Some (ValueOf uni) -> Term TyName name uni fun a
forall tyname name (uni :: * -> *) fun a.
a -> Some (ValueOf uni) -> Term tyname name uni fun a
Constant a
ann (Integer -> Some (ValueOf uni)
forall a (uni :: * -> *). Contains uni a => a -> Some (ValueOf uni)
PLC.someValue Integer
k))
        )