{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE Safe #-}
module ReWire.Fix (fix, fix', fixOn, boundedFixOn, fixOn', fixUntil, boundedFix, fixPure) where

import ReWire.Error (MonadError, AstError, failAt)
import ReWire.Annotation (noAnn)

import Control.Monad.Identity (Identity (runIdentity))
import Data.Hashable (Hashable (hash))
import Data.Text (Text)
import Numeric.Natural (Natural)

-- | Note: direct equality rather than comparing hashes: equality can stop at
--   the first difference, while hashing always traverses both terms fully.
fixPure :: Eq a => Natural -> (a -> a) -> a -> a
fixPure :: forall a. Eq a => Natural -> (a -> a) -> a -> a
fixPure Natural
n a -> a
f = Identity a -> a
forall a. Identity a -> a
runIdentity (Identity a -> a) -> (a -> Identity a) -> a -> a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (a -> a -> Bool) -> Natural -> (a -> Identity a) -> a -> Identity a
forall (m :: * -> *) a.
Monad m =>
(a -> a -> Bool) -> Natural -> (a -> m a) -> a -> m a
boundedFix a -> a -> Bool
forall a. Eq a => a -> a -> Bool
(==) Natural
n (a -> Identity a
forall a. a -> Identity a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (a -> Identity a) -> (a -> a) -> a -> Identity a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. a -> a
f)

fix :: (MonadError AstError m, Hashable a) => Text -> Natural -> (a -> m a) -> a -> m a
fix :: forall (m :: * -> *) a.
(MonadError AstError m, Hashable a) =>
Text -> Natural -> (a -> m a) -> a -> m a
fix = (a -> Int) -> Text -> Natural -> (a -> m a) -> a -> m a
forall (m :: * -> *) b a.
(MonadError AstError m, Eq b) =>
(a -> b) -> Text -> Natural -> (a -> m a) -> a -> m a
fixOn a -> Int
forall a. Hashable a => a -> Int
hash

fixUntil :: MonadError AstError m => (a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
fixUntil :: forall (m :: * -> *) a.
MonadError AstError m =>
(a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
fixUntil a -> Bool
h = (a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
forall (m :: * -> *) a.
MonadError AstError m =>
(a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
boundedFixOn ((a -> Bool) -> a -> a -> Bool
forall a b. a -> b -> a
const a -> Bool
h)

fixOn :: (MonadError AstError m, Eq b) => (a -> b) -> Text -> Natural -> (a -> m a) -> a -> m a
fixOn :: forall (m :: * -> *) b a.
(MonadError AstError m, Eq b) =>
(a -> b) -> Text -> Natural -> (a -> m a) -> a -> m a
fixOn a -> b
h = (a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
forall (m :: * -> *) a.
MonadError AstError m =>
(a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
boundedFixOn (\ a
a' a
a -> a -> b
h a
a' b -> b -> Bool
forall a. Eq a => a -> a -> Bool
== a -> b
h a
a)

-- | Note: evaluates `(f a)` at least once if the bound is not reached.
boundedFixOn :: MonadError AstError m => (a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
boundedFixOn :: forall (m :: * -> *) a.
MonadError AstError m =>
(a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
boundedFixOn a -> a -> Bool
_ Text
m Natural
0 a -> m a
_ a
_ = Annote -> Text -> m a
forall (m :: * -> *) an a.
(MonadError AstError m, Annotation an) =>
an -> Text -> m a
failAt Annote
noAnn (Text -> m a) -> Text -> m a
forall a b. (a -> b) -> a -> b
$ Text
m Text -> Text -> Text
forall a. Semigroup a => a -> a -> a
<> Text
" not terminating (mutually recursive definitions?)."
boundedFixOn a -> a -> Bool
h Text
m Natural
n a -> m a
f a
a = a -> m a
f a
a m a -> (a -> m a) -> m a
forall a b. m a -> (a -> m b) -> m b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= \ a
a' -> if a -> a -> Bool
h a
a' a
a then a -> m a
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure a
a' else (a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
forall (m :: * -> *) a.
MonadError AstError m =>
(a -> a -> Bool) -> Text -> Natural -> (a -> m a) -> a -> m a
boundedFixOn a -> a -> Bool
h Text
m (Natural
n Natural -> Natural -> Natural
forall a. Num a => a -> a -> a
- Natural
1) a -> m a
f a
a'

fix' :: Hashable a => (a -> a) -> a -> a
fix' :: forall a. Hashable a => (a -> a) -> a -> a
fix' = (a -> Int) -> (a -> a) -> a -> a
forall b a. Eq b => (a -> b) -> (a -> a) -> a -> a
fixOn' a -> Int
forall a. Hashable a => a -> Int
hash

-- | Note: the recursive call must reuse the already-forced @f a@ -- passing
--   the expression @f a@ itself builds an unshared thunk chain that gets
--   re-evaluated from scratch at every level, making the fixpoint
--   exponential in the iteration count.
fixOn' :: Eq b => (a -> b) -> (a -> a) -> a -> a
fixOn' :: forall b a. Eq b => (a -> b) -> (a -> a) -> a -> a
fixOn' a -> b
h a -> a
f a
a | a -> b
h a
a' b -> b -> Bool
forall a. Eq a => a -> a -> Bool
== a -> b
h a
a = a
a
             | Bool
otherwise   = (a -> b) -> (a -> a) -> a -> a
forall b a. Eq b => (a -> b) -> (a -> a) -> a -> a
fixOn' a -> b
h a -> a
f a
a'
      where a' :: a
a' = a -> a
f a
a

boundedFix :: Monad m => (a -> a -> Bool) -> Natural -> (a -> m a) -> a -> m a
boundedFix :: forall (m :: * -> *) a.
Monad m =>
(a -> a -> Bool) -> Natural -> (a -> m a) -> a -> m a
boundedFix a -> a -> Bool
_ Natural
0 a -> m a
_ a
a = a -> m a
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure a
a
boundedFix a -> a -> Bool
h Natural
n a -> m a
f a
a = a -> m a
f a
a m a -> (a -> m a) -> m a
forall a b. m a -> (a -> m b) -> m b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= \ a
a' -> if a -> a -> Bool
h a
a' a
a then a -> m a
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure a
a' else (a -> a -> Bool) -> Natural -> (a -> m a) -> a -> m a
forall (m :: * -> *) a.
Monad m =>
(a -> a -> Bool) -> Natural -> (a -> m a) -> a -> m a
boundedFix a -> a -> Bool
h (Natural
n Natural -> Natural -> Natural
forall a. Num a => a -> a -> a
- Natural
1) a -> m a
f a
a'