packages feed

proarrow-0.1.0.0: src/Proarrow/Promonad/State.hs

-- | The state promonad and its transformer: @'StateT' s p@ sandwiches @p@ between 'Reader' and 'Writer',
-- so @'State' s a b@ amounts to a map @s '**' a '~>' s '**' b@. It is only premonoidal, not monoidal:
-- the order in which two stateful effects run matters.
module Proarrow.Promonad.State where

import Prelude (($))

import Proarrow.Adjunction qualified as Adj
import Proarrow.Category.Instance.Opposite (OPPOSITE (..))
import Proarrow.Category.Instance.Prof (Prof (..))
import Proarrow.Category.Monoidal
  ( Monoidal (..)
  , MonoidalProfunctor (..)
  , SymMonoidal (..)
  , Tensor
  , leftUnitorWith
  , swap'
  )
import Proarrow.Category.Monoidal.Closed (Closed (..))
import Proarrow.Category.Monoidal.CompactClosed (CompactClosed (..))
import Proarrow.Category.Monoidal.Strength (Strong (..))
import Proarrow.Core (CategoryOf (..), Profunctor (..), Promonad (..), arr, obj, type (+->))
import Proarrow.Functor (Functor (..))
import Proarrow.Limit.BinaryProduct ((&&&))
import Proarrow.Monoid (Comonoid (..))
import Proarrow.Object (ObjDict (..), objDicts)
import Proarrow.Profunctor.Corepresentable (Corepresentable)
import Proarrow.Profunctor.Instance.Composition ((:.:) (..))
import Proarrow.Profunctor.Instance.Identity (Id (..))
import Proarrow.Profunctor.Representable (Representable (..))
import Proarrow.Promonad.Reader (Reader (..))
import Proarrow.Promonad.Writer (Writer (..))

type State s = StateT s Id
pattern State :: forall {k} a b s. (Monoidal k, Ob (s :: k)) => (Ob a, Ob b) => (s ** a) ~> (s ** b) -> State s a b
pattern State f <- (runStateT &&& objDicts -> (Id f, (ObjDict, ObjDict)))
  where
    State f = StateT (Reader id :.: Id f :.: Writer id) \\ f
{-# COMPLETE State #-}

-- | This is only premonoidal, not monoidal.
instance (SymMonoidal k, Ob s) => MonoidalProfunctor (State (s :: k)) where
  one = State (obj @s ** one) \\ (one :: (Unit :: k) ~> Unit)
  State @a1 @b1 f ** State @a2 @b2 g =
    let s = obj @s; a1 = obj @a1; b1 = obj @b1; a2 = obj @a2; b2 = obj @b2
    in State
         ( (s ** swap' b2 b1)
             . associator @_ @s @b2 @b1
             . (g ** b1)
             . associatorInv @_ @s @a2 @b1
             . (s ** swap' b1 a2)
             . associator @_ @s @b1 @a2
             . (f ** a2)
             . associatorInv @_ @s @a1 @a2
         )
         \\ (a1 ** a2)
         \\ (b1 ** b2)

type StateT :: k -> k +-> k -> k +-> k
newtype StateT s p a b where
  StateT :: (Reader (OP s) :.: p :.: Writer s) a b -> StateT s p a b

runStateT :: (Profunctor p) => StateT s p a b -> p (s ** a) (s ** b)
runStateT (StateT (Reader f :.: p :.: Writer g)) = dimap f g p

get :: forall {k} p s. (Promonad p, Monoidal k, Comonoid (s :: k)) => StateT s p Unit s
get = StateT (Reader (rightUnitor @k @s) :.: id :.: Writer (comult @s))

put :: forall {k} p s. (Promonad p, Monoidal k, Comonoid (s :: k)) => StateT s p s Unit
put = StateT (Reader (leftUnitorWith (counit @s)) :.: id :.: Writer (rightUnitorInv @k @s))

deriving newtype instance (Profunctor p, Monoidal k, Ob (s :: k)) => Profunctor (StateT s p)
deriving newtype instance (Representable p, Ob (s :: k), SymMonoidal k, Closed k) => Representable (StateT s p)
deriving newtype instance (Corepresentable p, Ob (s :: k), Monoidal k, CompactClosed k) => Corepresentable (StateT s p)

instance (Strong Tensor p, Ob (s :: k), SymMonoidal k) => Strong Tensor (StateT s p) where
  act @a (StateT p) = StateT (act @Tensor @_ @a p)

instance (Ob (s :: k), Monoidal k, Strong Tensor p, Promonad p) => Promonad (StateT s p) where
  id @a = withOb2 @k @s @a $ StateT (Reader id :.: act @Tensor @p @s (id @p @a) :.: Writer id)
  StateT (r1 :.: p1 :.: w1) . StateT (r2 :.: p2 :.: w2) = StateT (r2 :.: (p1 . arr (Adj.counit (w2 :.: r1)) . p2) :.: w1)

instance (Monoidal k, Ob s) => Functor (StateT s :: k +-> k -> k +-> k) where
  map (Prof n) = Prof \(StateT (r :.: p :.: w)) -> StateT (r :.: n p :.: w)