packages feed

imsos-monad-0.2.4.0: src/Control/Monad/IMSOS/Monad.hs

{-# LANGUAGE UndecidableInstances
  , FlexibleInstances
  , FlexibleContexts
  , MultiParamTypeClasses
  , TypeOperators
  , LambdaCase
  , ScopedTypeVariables
  , TypeApplications
  , TupleSections
  , AllowAmbiguousTypes
#-}

module Control.Monad.IMSOS.Monad
  (MonadIMSOS(..), runIMSOS, yieldIMSOS
  ,Control.Monad.IMSOS.Monad.tell, Control.Monad.IMSOS.Monad.listen, Control.Monad.IMSOS.Monad.censor
  ,Control.Monad.IMSOS.Monad.local, Control.Monad.IMSOS.Monad.reader, Control.Monad.IMSOS.Monad.ask
  ,Control.Monad.IMSOS.Monad.get, Control.Monad.IMSOS.Monad.put, Control.Monad.IMSOS.Monad.modify
  ,fail, throwError, catchError
  ,guard, (<|>)
  ) where

import Control.Applicative (Alternative(..))
import Control.Monad.Error.Class
import Control.Monad.Reader (MonadReader(..))
import Control.Monad.State  (MonadState(..), modify)
import Control.Monad.Writer (MonadWriter(..), censor)
import Control.Monad (ap, join, guard)

import Data.Foldable (toList)
import Data.Comp.ProjectionExt ((:*:)(..), (:<~), (:<|), pr, uncons, recons, modify)

newtype MonadIMSOS r s w me e fl a = MonadIMSOS {
    unwrap :: r -> s -> me e (fl (a, s, w))
  }

instance (Functor (me e), Functor fl) => Functor (MonadIMSOS r s w me e fl) where
  fmap f (MonadIMSOS m) = MonadIMSOS m'
    where m' r s = fmap modify' <$> m r s
           where modify' (a, s', w) = (f a, s', w)

instance (Monad (me e), Monoid w, Traversable fl, Monad fl)
    => Monad (MonadIMSOS r s w me e fl) where
  (MonadIMSOS p) >>= m = MonadIMSOS $ \r1 s1  -> do
     p_res <- p r1 s1
     let modify' (a, s2, w2) = fmap (,w2) <$> q r1 s2
          where MonadIMSOS q = m a
     q_ress <- mapM modify' p_res
     let mod2 :: ((b, s, w), w) -> (b, s, w)
         mod2 ((b, s3, w3), w2) = (b, s3, w2 <> w3)
     return (fmap mod2 (join q_ress))

instance (Monad (me e), Monoid w, Monad fl, Traversable fl)
    => Applicative (MonadIMSOS r s w me e fl) where
  pure a = MonadIMSOS m
    where m _ s = pure (pure (a, s, mempty))
  (<*>) = ap

instance (Monad (me e), MonadError e (me e), Monoid w
         ,Alternative fl, Monad fl, Traversable fl)
    => Alternative (MonadIMSOS r s w me e fl) where
  empty = MonadIMSOS (\_ _ -> return empty)
  (MonadIMSOS p) <|> (MonadIMSOS q) = MonadIMSOS m
    where m r s = tryError (p r s) >>= \case
                      Left _      -> q r s
                      Right p_res -> tryError (q r s) >>= \case
                          Left _ -> return empty
                          Right q_res -> return (p_res <|> q_res)

instance {-# INCOHERENT #-}
  (MonadError String (me String), Monad (me String), Monoid w, Traversable fl, Monad fl)
    => MonadFail (MonadIMSOS r s w me String fl) where
  fail = throwError

instance (MonadError e (me e), Monad (me e), Monoid w, Monad fl, Traversable fl)
    => MonadReader r (MonadIMSOS r s w me e fl) where
  ask = MonadIMSOS m
    where m r s = pure $ pure (r, s, mempty)
  local f (MonadIMSOS p) = MonadIMSOS m
    where m r s = p (f r) s

instance (MonadError e (me e), Monad (me e), Monoid w, Monad fl, Traversable fl)
    => MonadState s (MonadIMSOS r s w me e fl) where
  state act = MonadIMSOS m
    where m _ s = pure $ pure (a, s', mempty)
           where (a, s') = act s

instance (MonadError e (me e), Monad (me e), Monoid w, Monad fl, Traversable fl)
    => MonadWriter w (MonadIMSOS r s w me e fl) where
  tell w = MonadIMSOS m
    where m _ s = pure $ pure ((), s, w)
  listen (MonadIMSOS p) = MonadIMSOS m
    where m r s = p r s >>= \as -> pure (fmap modify' as)
            where modify' (a, s', w) = ((a, w), s', w)
  pass (MonadIMSOS p) = MonadIMSOS m
    where m r s = p r s >>= \as -> pure (fmap modify' as)
            where modify' ((a,f), s', w) = (a, s', f w)

instance (MonadError e (me e), Monad (me e), Monoid w, Traversable fl, Monad fl)
    => MonadError e (MonadIMSOS r s w me e fl) where
  throwError e = MonadIMSOS m
    where m _ _ = throwError e
  catchError (MonadIMSOS p) h = MonadIMSOS m
    where m r s = p r s `catchError` \e ->
                  let MonadIMSOS q = h e
                  in q r s -- continues with a state as if p not executed

class Unifies a where
    unify :: [a] -> a

search :: (Unifies s, Unifies w, Monad (me e), Applicative fl, Foldable fl)
  => MonadIMSOS r s w me e fl a -> MonadIMSOS r s w me e fl [a]
search (MonadIMSOS p) = MonadIMSOS $ \r s -> do
  (as, ss, ws) <- unzip3 . toList <$> p r s
  return $ pure (as, unify ss, unify ws)

runIMSOS :: (MonadError e (me e))
  => r -> s -> MonadIMSOS r s w me e fl a -> me e (fl (a, s, w))
runIMSOS r s (MonadIMSOS f) = f r s

yieldIMSOS :: (Functor (me e), Functor fl, MonadFail (me e), MonadError e (me e)) =>
  r -> s -> ((a, s, w) -> y) -> MonadIMSOS r s w me e fl a -> me e (fl y)
yieldIMSOS r s toyield m = fmap toyield <$> runIMSOS r s m

-- ------------------------------
-- --- Auxiliary entities 
-- ------------------------------

instance (Semigroup f, Semigroup g)
  => Semigroup (f :*: g) where
  (l1 :*: r1) <> (l2 :*: r2) = (l1 <> l2) :*: (r1 <> r2)

instance (Monoid f, Monoid g)
  => Monoid (f :*: g) where
    mempty = mempty :*: mempty

tell :: forall f w m. (f :<| w, MonadWriter w m) => f -> m ()
tell f = Control.Monad.Writer.tell (recons rem_ f)
  where (_, rem_) = uncons @f (mempty @w)

listen :: forall f w m a. (f :<~ w, MonadWriter w m) => m a -> m (a, f)
listen m = fmap (pr @f) <$> Control.Monad.Writer.listen m

censor :: forall f w m a. (f :<| w, MonadWriter w m) => (f -> f) -> m a ->  m a
censor tr = Control.Monad.Writer.censor (Data.Comp.ProjectionExt.modify @f tr)

get :: forall f s m. (f :<~ s, MonadState s m) => m f
get = pr @f <$> Control.Monad.State.get

put :: forall f s m. (f :<| s, MonadState s m) => f -> m ()
put f = Control.Monad.IMSOS.Monad.modify (const f)

modify :: forall f s m. (f :<| s, MonadState s m) => (f -> f) -> m ()
modify tr = Control.Monad.State.modify (Data.Comp.ProjectionExt.modify @f tr)

ask :: forall f r m. (f :<~ r, MonadReader r m) => m f
ask = pr @f <$> Control.Monad.Reader.ask

reader :: forall f r m a. (f :<| r, MonadReader r m) => (f -> a) -> m a
reader tr = Control.Monad.Reader.reader tr'
  where tr' r = tr (pr @f r)

local :: forall f r m a. (f :<| r, MonadReader r m) => (f -> f) -> m a -> m a
local tr = Control.Monad.Reader.local (Data.Comp.ProjectionExt.modify tr)