-- | Batch finalization state and ordered finalization operations.
--
-- A finalization commits the outcome, records it, and then calls its hooks. The
-- record precedes the hooks. 'settle' enforces this order.
module Arbiter.Worker.Settle
  ( -- * Handoff
    CancelHandoff
  , newCancelHandoff
  , pendingJobs
  , unownedJobs
  , recordCancelled
  , cancelFinalized
  , markCancelFinalized

    -- * Settling
  , Settled
  , finalized
  , disowned
  , settle
  , settleBy
  , settleInterruptibly
  , record

    -- * Job sets
  , hasIdIn
  , byIdDesc
  ) where

import Arbiter.Core.Job.Types qualified as Job
import Control.Monad.IO.Class (MonadIO, liftIO)
import Data.Foldable (toList)
import Data.IORef (IORef, atomicModifyIORef', newIORef, readIORef)
import Data.Int (Int64)
import Data.List (sortOn)
import Data.List.NonEmpty (NonEmpty)
import Data.Ord (Down (..))
import Data.Set qualified as Set
import UnliftIO (MonadUnliftIO, mask_)

-- | What a batch has settled so far.
data BatchProgress = BatchProgress
  { BatchProgress -> Set Int64
progressHandled :: !(Set.Set Int64)
  -- ^ Jobs whose outcome has been recorded.
  , BatchProgress -> Set Int64
progressUnowned :: !(Set.Set Int64)
  -- ^ Jobs a batch ack found under another claim.
  , BatchProgress -> Set Int64
progressCancelled :: !(Set.Set Int64)
  -- ^ Jobs a force-cancel accounted for.
  , BatchProgress -> Bool
progressFinalized :: !Bool
  -- ^ Whether the force-cancel finalizer ran to completion.
  }

-- | Finalization state shared by the handler and force-cancel operation.
newtype CancelHandoff = CancelHandoff (IORef BatchProgress)

newCancelHandoff :: (MonadIO m) => m CancelHandoff
newCancelHandoff :: forall (m :: * -> *). MonadIO m => m CancelHandoff
newCancelHandoff = IO CancelHandoff -> m CancelHandoff
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IORef BatchProgress -> CancelHandoff
CancelHandoff (IORef BatchProgress -> CancelHandoff)
-> IO (IORef BatchProgress) -> IO CancelHandoff
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> BatchProgress -> IO (IORef BatchProgress)
forall a. a -> IO (IORef a)
newIORef (Set Int64 -> Set Int64 -> Set Int64 -> Bool -> BatchProgress
BatchProgress Set Int64
forall a. Monoid a => a
mempty Set Int64
forall a. Monoid a => a
mempty Set Int64
forall a. Monoid a => a
mempty Bool
False))

readProgress :: (MonadIO m) => CancelHandoff -> m BatchProgress
readProgress :: forall (m :: * -> *). MonadIO m => CancelHandoff -> m BatchProgress
readProgress (CancelHandoff IORef BatchProgress
ref) = IO BatchProgress -> m BatchProgress
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IORef BatchProgress -> IO BatchProgress
forall a. IORef a -> IO a
readIORef IORef BatchProgress
ref)

onProgress :: (MonadIO m) => CancelHandoff -> (BatchProgress -> (BatchProgress, a)) -> m a
onProgress :: forall (m :: * -> *) a.
MonadIO m =>
CancelHandoff -> (BatchProgress -> (BatchProgress, a)) -> m a
onProgress (CancelHandoff IORef BatchProgress
ref) = IO a -> m a
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO a -> m a)
-> ((BatchProgress -> (BatchProgress, a)) -> IO a)
-> (BatchProgress -> (BatchProgress, a))
-> m a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. IORef BatchProgress
-> (BatchProgress -> (BatchProgress, a)) -> IO a
forall a b. IORef a -> (a -> (a, b)) -> IO b
atomicModifyIORef' IORef BatchProgress
ref

-- | Whether a set of job ids names this job.
hasIdIn :: Set.Set Int64 -> Job.JobRead payload -> Bool
hasIdIn :: forall payload. Set Int64 -> JobRead payload -> Bool
hasIdIn Set Int64
ids JobRead payload
job = Int64 -> Set Int64 -> Bool
forall a. Ord a => a -> Set a -> Bool
Set.member (JobRead payload -> Int64
forall payload key q insertedAt adm.
JobRecord payload key q insertedAt adm -> key
Job.primaryKey JobRead payload
job) Set Int64
ids

-- | Select jobs from the batch in descending identifier order. This gives ack,
-- force-cancel, and heartbeat operations the same row-lock order.
byIdDesc :: (Job.JobRead payload -> Bool) -> NonEmpty (Job.JobRead payload) -> [Job.JobRead payload]
byIdDesc :: forall payload.
(JobRead payload -> Bool)
-> NonEmpty (JobRead payload) -> [JobRead payload]
byIdDesc JobRead payload -> Bool
keep = (JobRead payload -> Down Int64)
-> [JobRead payload] -> [JobRead payload]
forall b a. Ord b => (a -> b) -> [a] -> [a]
sortOn (Int64 -> Down Int64
forall a. a -> Down a
Down (Int64 -> Down Int64)
-> (JobRead payload -> Int64) -> JobRead payload -> Down Int64
forall b c a. (b -> c) -> (a -> b) -> a -> c
. JobRead payload -> Int64
forall payload key q insertedAt adm.
JobRecord payload key q insertedAt adm -> key
Job.primaryKey) ([JobRead payload] -> [JobRead payload])
-> (NonEmpty (JobRead payload) -> [JobRead payload])
-> NonEmpty (JobRead payload)
-> [JobRead payload]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (JobRead payload -> Bool) -> [JobRead payload] -> [JobRead payload]
forall a. (a -> Bool) -> [a] -> [a]
filter JobRead payload -> Bool
keep ([JobRead payload] -> [JobRead payload])
-> (NonEmpty (JobRead payload) -> [JobRead payload])
-> NonEmpty (JobRead payload)
-> [JobRead payload]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. NonEmpty (JobRead payload) -> [JobRead payload]
forall a. NonEmpty a -> [a]
forall (t :: * -> *) a. Foldable t => t a -> [a]
toList

-- | Jobs in the batch that have no recorded outcome.
pendingJobs :: (MonadIO m) => CancelHandoff -> NonEmpty (Job.JobRead payload) -> m [Job.JobRead payload]
pendingJobs :: forall (m :: * -> *) payload.
MonadIO m =>
CancelHandoff -> NonEmpty (JobRead payload) -> m [JobRead payload]
pendingJobs CancelHandoff
handoff NonEmpty (JobRead payload)
jobs = do
  progress <- CancelHandoff -> m BatchProgress
forall (m :: * -> *). MonadIO m => CancelHandoff -> m BatchProgress
readProgress CancelHandoff
handoff
  pure (byIdDesc (not . hasIdIn (progressHandled progress)) jobs)

-- | The batch's jobs a settle found under another claim.
unownedJobs :: (MonadIO m) => CancelHandoff -> NonEmpty (Job.JobRead payload) -> m [Job.JobRead payload]
unownedJobs :: forall (m :: * -> *) payload.
MonadIO m =>
CancelHandoff -> NonEmpty (JobRead payload) -> m [JobRead payload]
unownedJobs CancelHandoff
handoff NonEmpty (JobRead payload)
jobs = do
  progress <- CancelHandoff -> m BatchProgress
forall (m :: * -> *). MonadIO m => CancelHandoff -> m BatchProgress
readProgress CancelHandoff
handoff
  pure (byIdDesc (hasIdIn (progressUnowned progress)) jobs)

-- | Add to the jobs a force-cancel accounted for, returning the ids new to this call
-- and every id recorded so far.
recordCancelled :: (MonadIO m) => CancelHandoff -> Set.Set Int64 -> m (Set.Set Int64, Set.Set Int64)
recordCancelled :: forall (m :: * -> *).
MonadIO m =>
CancelHandoff -> Set Int64 -> m (Set Int64, Set Int64)
recordCancelled CancelHandoff
handoff Set Int64
ids =
  CancelHandoff
-> (BatchProgress -> (BatchProgress, (Set Int64, Set Int64)))
-> m (Set Int64, Set Int64)
forall (m :: * -> *) a.
MonadIO m =>
CancelHandoff -> (BatchProgress -> (BatchProgress, a)) -> m a
onProgress CancelHandoff
handoff ((BatchProgress -> (BatchProgress, (Set Int64, Set Int64)))
 -> m (Set Int64, Set Int64))
-> (BatchProgress -> (BatchProgress, (Set Int64, Set Int64)))
-> m (Set Int64, Set Int64)
forall a b. (a -> b) -> a -> b
$ \BatchProgress
progress ->
    let cancelled :: Set Int64
cancelled = BatchProgress -> Set Int64
progressCancelled BatchProgress
progress Set Int64 -> Set Int64 -> Set Int64
forall a. Semigroup a => a -> a -> a
<> Set Int64
ids
     in (BatchProgress
progress {progressCancelled = cancelled}, (Set Int64
ids Set Int64 -> Set Int64 -> Set Int64
forall a. Ord a => Set a -> Set a -> Set a
Set.\\ BatchProgress -> Set Int64
progressCancelled BatchProgress
progress, Set Int64
cancelled))

-- | Whether the force-cancel finalizer ran to completion.
cancelFinalized :: (MonadIO m) => CancelHandoff -> m Bool
cancelFinalized :: forall (m :: * -> *). MonadIO m => CancelHandoff -> m Bool
cancelFinalized = (BatchProgress -> Bool) -> m BatchProgress -> m Bool
forall a b. (a -> b) -> m a -> m b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap BatchProgress -> Bool
progressFinalized (m BatchProgress -> m Bool)
-> (CancelHandoff -> m BatchProgress) -> CancelHandoff -> m Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. CancelHandoff -> m BatchProgress
forall (m :: * -> *). MonadIO m => CancelHandoff -> m BatchProgress
readProgress

markCancelFinalized :: (MonadIO m) => CancelHandoff -> m ()
markCancelFinalized :: forall (m :: * -> *). MonadIO m => CancelHandoff -> m ()
markCancelFinalized CancelHandoff
handoff = CancelHandoff -> (BatchProgress -> (BatchProgress, ())) -> m ()
forall (m :: * -> *) a.
MonadIO m =>
CancelHandoff -> (BatchProgress -> (BatchProgress, a)) -> m a
onProgress CancelHandoff
handoff ((BatchProgress -> (BatchProgress, ())) -> m ())
-> (BatchProgress -> (BatchProgress, ())) -> m ()
forall a b. (a -> b) -> a -> b
$ \BatchProgress
progress -> (BatchProgress
progress {progressFinalized = True}, ())

-- | What a settle accounted for: the jobs it finalized, and the jobs it found under
-- another claim.
data Settled payload = Settled [Job.JobRead payload] [Job.JobRead payload]

instance Semigroup (Settled payload) where
  Settled [JobRead payload]
handledA [JobRead payload]
unownedA <> :: Settled payload -> Settled payload -> Settled payload
<> Settled [JobRead payload]
handledB [JobRead payload]
unownedB = [JobRead payload] -> [JobRead payload] -> Settled payload
forall payload.
[JobRead payload] -> [JobRead payload] -> Settled payload
Settled ([JobRead payload]
handledA [JobRead payload] -> [JobRead payload] -> [JobRead payload]
forall a. Semigroup a => a -> a -> a
<> [JobRead payload]
handledB) ([JobRead payload]
unownedA [JobRead payload] -> [JobRead payload] -> [JobRead payload]
forall a. Semigroup a => a -> a -> a
<> [JobRead payload]
unownedB)

-- | Jobs a settle finalized.
finalized :: [Job.JobRead payload] -> Settled payload
finalized :: forall payload. [JobRead payload] -> Settled payload
finalized [JobRead payload]
jobs = [JobRead payload] -> [JobRead payload] -> Settled payload
forall payload.
[JobRead payload] -> [JobRead payload] -> Settled payload
Settled [JobRead payload]
jobs []

-- | Jobs a settle found under another claim.
disowned :: [Job.JobRead payload] -> Settled payload
disowned :: forall payload. [JobRead payload] -> Settled payload
disowned [JobRead payload]
jobs = [JobRead payload] -> [JobRead payload] -> Settled payload
forall payload.
[JobRead payload] -> [JobRead payload] -> Settled payload
Settled [] [JobRead payload]
jobs

-- | Record a finalization in one atomic update. Jobs owned by another claim also
-- count as handled.
record :: (MonadIO m) => CancelHandoff -> Settled payload -> m ()
record :: forall (m :: * -> *) payload.
MonadIO m =>
CancelHandoff -> Settled payload -> m ()
record CancelHandoff
handoff (Settled [JobRead payload]
handled [JobRead payload]
unowned) =
  CancelHandoff -> (BatchProgress -> (BatchProgress, ())) -> m ()
forall (m :: * -> *) a.
MonadIO m =>
CancelHandoff -> (BatchProgress -> (BatchProgress, a)) -> m a
onProgress CancelHandoff
handoff ((BatchProgress -> (BatchProgress, ())) -> m ())
-> (BatchProgress -> (BatchProgress, ())) -> m ()
forall a b. (a -> b) -> a -> b
$ \BatchProgress
progress ->
    ( BatchProgress
progress
        { progressHandled = progressHandled progress <> ids handled <> gone
        , progressUnowned = progressUnowned progress <> gone
        }
    , ()
    )
  where
    gone :: Set Int64
gone = [JobRead payload] -> Set Int64
forall {payload} {q} {insertedAt} {adm}.
[JobRecord payload Int64 q insertedAt adm] -> Set Int64
ids [JobRead payload]
unowned
    ids :: [JobRecord payload Int64 q insertedAt adm] -> Set Int64
ids = [Int64] -> Set Int64
forall a. Ord a => [a] -> Set a
Set.fromList ([Int64] -> Set Int64)
-> ([JobRecord payload Int64 q insertedAt adm] -> [Int64])
-> [JobRecord payload Int64 q insertedAt adm]
-> Set Int64
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (JobRecord payload Int64 q insertedAt adm -> Int64)
-> [JobRecord payload Int64 q insertedAt adm] -> [Int64]
forall a b. (a -> b) -> [a] -> [b]
map JobRecord payload Int64 q insertedAt adm -> Int64
forall payload key q insertedAt adm.
JobRecord payload key q insertedAt adm -> key
Job.primaryKey

-- | Commit a settle, record it, then run its hooks. @protect@ covers the commit and
-- the record together.
settleWith
  :: (MonadIO m)
  => (m a -> m a)
  -> CancelHandoff
  -> m a
  -- ^ Commit.
  -> (a -> Settled payload)
  -- ^ What it accounted for.
  -> (a -> m b)
  -- ^ Hooks. The record precedes them.
  -> m b
settleWith :: forall (m :: * -> *) a payload b.
MonadIO m =>
(m a -> m a)
-> CancelHandoff
-> m a
-> (a -> Settled payload)
-> (a -> m b)
-> m b
settleWith m a -> m a
protect CancelHandoff
handoff m a
commit a -> Settled payload
accounts a -> m b
hooks = do
  outcome <- m a -> m a
protect (m a -> m a) -> m a -> m a
forall a b. (a -> b) -> a -> b
$ m a
commit 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
result -> a
result a -> m () -> m a
forall a b. a -> m b -> m a
forall (f :: * -> *) a b. Functor f => a -> f b -> f a
<$ CancelHandoff -> Settled payload -> m ()
forall (m :: * -> *) payload.
MonadIO m =>
CancelHandoff -> Settled payload -> m ()
record CancelHandoff
handoff (a -> Settled payload
accounts a
result)
  hooks outcome

-- | Apply 'settleWith' to a set known before the commit. Mask asynchronous
-- exceptions between the commit and state update.
settle
  :: (MonadUnliftIO m)
  => CancelHandoff
  -> Settled payload
  -> m a
  -> (a -> m b)
  -> m b
settle :: forall (m :: * -> *) payload a b.
MonadUnliftIO m =>
CancelHandoff -> Settled payload -> m a -> (a -> m b) -> m b
settle CancelHandoff
handoff Settled payload
settled m a
commit = (m a -> m a)
-> CancelHandoff
-> m a
-> (a -> Settled payload)
-> (a -> m b)
-> m b
forall (m :: * -> *) a payload b.
MonadIO m =>
(m a -> m a)
-> CancelHandoff
-> m a
-> (a -> Settled payload)
-> (a -> m b)
-> m b
settleWith m a -> m a
forall (m :: * -> *) a. MonadUnliftIO m => m a -> m a
mask_ CancelHandoff
handoff m a
commit (Settled payload -> a -> Settled payload
forall a b. a -> b -> a
const Settled payload
settled)

-- | 'settle' for a commit that decides what it settled.
settleBy
  :: (MonadUnliftIO m)
  => CancelHandoff
  -> m a
  -> (a -> Settled payload)
  -> (a -> m b)
  -> m b
settleBy :: forall (m :: * -> *) a payload b.
MonadUnliftIO m =>
CancelHandoff -> m a -> (a -> Settled payload) -> (a -> m b) -> m b
settleBy = (m a -> m a)
-> CancelHandoff
-> m a
-> (a -> Settled payload)
-> (a -> m b)
-> m b
forall (m :: * -> *) a payload b.
MonadIO m =>
(m a -> m a)
-> CancelHandoff
-> m a
-> (a -> Settled payload)
-> (a -> m b)
-> m b
settleWith m a -> m a
forall (m :: * -> *) a. MonadUnliftIO m => m a -> m a
mask_

-- | Apply 'settle' with an interruptible commit for a transaction that contains
-- the handler.
settleInterruptibly
  :: (MonadIO m)
  => CancelHandoff
  -> Settled payload
  -> m a
  -> (a -> m b)
  -> m b
settleInterruptibly :: forall (m :: * -> *) payload a b.
MonadIO m =>
CancelHandoff -> Settled payload -> m a -> (a -> m b) -> m b
settleInterruptibly CancelHandoff
handoff Settled payload
settled m a
commit = (m a -> m a)
-> CancelHandoff
-> m a
-> (a -> Settled payload)
-> (a -> m b)
-> m b
forall (m :: * -> *) a payload b.
MonadIO m =>
(m a -> m a)
-> CancelHandoff
-> m a
-> (a -> Settled payload)
-> (a -> m b)
-> m b
settleWith m a -> m a
forall a. a -> a
id CancelHandoff
handoff m a
commit (Settled payload -> a -> Settled payload
forall a b. a -> b -> a
const Settled payload
settled)