{-# LANGUAGE OverloadedStrings #-}

module Arbiter.Worker.Dispatcher
  ( runDispatcher
  ) where

import Arbiter.Core.HighLevel (QueueOperation)
import Arbiter.Core.HighLevel qualified as Arb
import Arbiter.Core.Job.Types (JobRead)
import Arbiter.Core.Listen (Notification)
import Arbiter.Core.Operations qualified as Ops
import Control.Monad (void)
import Data.Foldable (traverse_)
import Data.List.NonEmpty (NonEmpty (..))
import UnliftIO.STM qualified as STM

import Arbiter.Worker.Config
  ( HandlerMode (..)
  , WorkerConfig (..)
  , handlerBatchSize
  , heartbeatSignal
  , readEffectiveState
  )
import Arbiter.Worker.Logger (LogLevel (..), newFailureGate, tryReported)
import Arbiter.Worker.NotificationListener (runNotificationConsumer)

-- | Wake on NOTIFY, poll timer, or worker-finished, then claim up to capacity.
-- @notifVar@ is filled from the shared hub in "Arbiter.Core.Listen".
runDispatcher
  :: forall payload m
   . (QueueOperation m payload)
  => WorkerConfig m payload
  -> Int
  -> STM.TBQueue (NonEmpty (JobRead payload))
  -> STM.TVar Int
  -> STM.TVar Bool
  -> STM.TVar (Maybe Notification)
  -> m ()
runDispatcher :: forall payload (m :: * -> *).
QueueOperation m payload =>
WorkerConfig m payload
-> Int
-> TBQueue (NonEmpty (JobRead payload))
-> TVar Int
-> TVar Bool
-> TVar (Maybe Notification)
-> m ()
runDispatcher WorkerConfig m payload
config Int
workerCapacity TBQueue (NonEmpty (JobRead payload))
workQueue TVar Int
busyWorkerCount TVar Bool
workerFinishedVar TVar (Maybe Notification)
notifVar = do
  -- The claim statement varies only with free capacity. Render every variant once.
  claimSql <-
    forall payload (m :: * -> *).
QueueOperation m payload =>
Int -> Int -> NominalDiffTime -> UUID -> m ClaimSql
Arb.mkClaimSql @payload (WorkerConfig m payload -> Int
forall (m :: * -> *) payload. WorkerConfig m payload -> Int
handlerBatchSize WorkerConfig m payload
config) Int
workerCapacity (WorkerConfig m payload -> NominalDiffTime
forall (m :: * -> *) payload.
WorkerConfig m payload -> NominalDiffTime
visibilityTimeout WorkerConfig m payload
config) (WorkerConfig m payload -> UUID
forall (m :: * -> *) payload. WorkerConfig m payload -> UUID
workerId WorkerConfig m payload
config)
  claimGate <- newFailureGate
  let
    calcFreeWorkers :: STM.STM Int
    calcFreeWorkers = do
      busyCount <- TVar Int -> STM Int
forall a. TVar a -> STM a
STM.readTVar TVar Int
busyWorkerCount
      queuedCount <- fromIntegral <$> STM.lengthTBQueue workQueue
      pure $ workerCapacity - (busyCount + queuedCount)

    getFreeWorkers :: STM.STM (Maybe Int)
    getFreeWorkers = do
      free <- STM Int
calcFreeWorkers
      pure $ if free > 0 then Just free else Nothing

    claimAndEnqueue :: Int -> m ()
    claimAndEnqueue Int
freeWorkers = do
      eJobs <- LogConfig
-> LogLevel
-> FailureGate
-> Text
-> m [NonEmpty (JobRead payload)]
-> m (Either SomeException [NonEmpty (JobRead payload)])
forall (m :: * -> *) a.
MonadUnliftIO m =>
LogConfig
-> LogLevel
-> FailureGate
-> Text
-> m a
-> m (Either SomeException a)
tryReported (WorkerConfig m payload -> LogConfig
forall (m :: * -> *) payload. WorkerConfig m payload -> LogConfig
logConfig WorkerConfig m payload
config) LogLevel
Error FailureGate
claimGate Text
"Dispatcher claim" (m [NonEmpty (JobRead payload)]
 -> m (Either SomeException [NonEmpty (JobRead payload)]))
-> m [NonEmpty (JobRead payload)]
-> m (Either SomeException [NonEmpty (JobRead payload)])
forall a b. (a -> b) -> a -> b
$
        case WorkerConfig m payload -> HandlerMode m payload
forall (m :: * -> *) payload.
WorkerConfig m payload -> HandlerMode m payload
handlerMode WorkerConfig m payload
config of
          SingleJobMode JobHandler m payload (ResultOf m payload)
_ ->
            (JobRead payload -> NonEmpty (JobRead payload))
-> [JobRead payload] -> [NonEmpty (JobRead payload)]
forall a b. (a -> b) -> [a] -> [b]
map (JobRead payload -> [JobRead payload] -> NonEmpty (JobRead payload)
forall a. a -> [a] -> NonEmpty a
:| []) ([JobRead payload] -> [NonEmpty (JobRead payload)])
-> m [JobRead payload] -> m [NonEmpty (JobRead payload)]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> ClaimSql -> Int -> m [JobRead payload]
forall (m :: * -> *) payload.
(JobPayload payload, MonadArbiter m) =>
ClaimSql -> Int -> m [JobRead payload]
Ops.claimJobsCached ClaimSql
claimSql Int
freeWorkers
          BatchedJobsMode Int
_ NonEmpty (JobRead payload)
-> BatchCallbacks m payload (ResultOf m payload) -> m ()
_ ->
            ClaimSql -> Int -> m [NonEmpty (JobRead payload)]
forall (m :: * -> *) payload.
(JobPayload payload, MonadArbiter m) =>
ClaimSql -> Int -> m [NonEmpty (JobRead payload)]
Ops.claimJobsBatchedCached ClaimSql
claimSql Int
freeWorkers
      traverse_ (STM.atomically . traverse_ (STM.writeTBQueue workQueue)) eJobs
      -- Pulse on every attempt, including a failed claim.
      STM.atomically $ void $ STM.tryPutTMVar (heartbeatSignal config) ()

    claimOnWakeup :: m ()
    claimOnWakeup = do
      mFree <- STM (Maybe Int) -> m (Maybe Int)
forall (m :: * -> *) a. MonadIO m => STM a -> m a
STM.atomically STM (Maybe Int)
getFreeWorkers
      traverse_ claimAndEnqueue mFree

    workerFinishedTrigger = STM () -> Maybe (STM ())
forall a. a -> Maybe a
Just (STM () -> Maybe (STM ())) -> STM () -> Maybe (STM ())
forall a b. (a -> b) -> a -> b
$ do
      finished <- TVar Bool -> STM Bool
forall a. TVar a -> STM a
STM.readTVar TVar Bool
workerFinishedVar
      STM.checkSTM finished
      STM.writeTVar workerFinishedVar False

  runNotificationConsumer
    (readEffectiveState config)
    (pollInterval config)
    notifVar
    workerFinishedTrigger
    (const claimOnWakeup)