{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeFamilies #-}

-- | Handler execution and job settlement for a worker pool.
module Arbiter.Worker.Processing
  ( workerLoop
  ) where

import Arbiter.Core.Exceptions (JobForceCancelled (..), displayEx)
import Arbiter.Core.HighLevel (JobOperation)
import Arbiter.Core.HighLevel qualified as Arb
import Arbiter.Core.Job.Types qualified as Job
import Arbiter.Core.JobResult
import Arbiter.Core.MonadArbiter (MonadArbiter (..))
import Arbiter.Core.Trace (ConsumeSpan, resolveTracer, withConsumeSpan)
import Control.Exception (fromException)
import Control.Exception qualified as E
import Control.Monad (forever, unless)
import Control.Monad.IO.Class (liftIO)
import Data.Foldable (toList, traverse_)
import Data.List.NonEmpty (NonEmpty (..))
import Data.Time (getCurrentTime)
import UnliftIO
  ( atomically
  , catchSyncOrAsync
  , finally
  , mask_
  , modifyTVar'
  , readTBQueue
  , tryAny
  , writeTVar
  )
import UnliftIO.Async qualified as Async
import UnliftIO.Concurrent (threadDelay)
import UnliftIO.STM (TBQueue, TVar)

import Arbiter.Worker.ChannelHandlers (RunningJobs, withRegisteredJobs)
import Arbiter.Worker.Config
import Arbiter.Worker.Heartbeat (withJobsHeartbeat)
import Arbiter.Worker.Logger
import Arbiter.Worker.Logger.Internal (runHook)
import Arbiter.Worker.Results (storeJobResult)
import Arbiter.Worker.Settle
  ( CancelHandoff
  , cancelFinalized
  , finalized
  , markCancelFinalized
  , newCancelHandoff
  , pendingJobs
  , settleInterruptibly
  )
import Arbiter.Worker.Settlement
  ( ackOrGone
  , batchCallbacks
  , batchLog
  , finalizeForceCancelled
  , jobLog
  , reportBatchOutcome
  , reportSuccess
  )

-- | Main loop for a single worker thread.
workerLoop
  :: forall payload m
   . ( EncodeJobResult (ResultOf m payload)
     , JobOperation m payload
     )
  => WorkerConfig m payload
  -> ConsumeSpan
  -- ^ The pool's consumer-span shape, built once for its queue.
  -> RunningJobs
  -- ^ Pool-shared map from job id to running handler async.
  -> TBQueue (NonEmpty (Job.JobRead payload))
  -> TVar Int
  -- ^ Busy worker count
  -> TVar Bool
  -- ^ Worker finished signal
  -> m ()
workerLoop :: forall payload (m :: * -> *).
(EncodeJobResult (ResultOf m payload), JobOperation m payload) =>
WorkerConfig m payload
-> ConsumeSpan
-> RunningJobs
-> TBQueue (NonEmpty (JobRead payload))
-> TVar Int
-> TVar Bool
-> m ()
workerLoop WorkerConfig m payload
config ConsumeSpan
consumeSpan RunningJobs
runningJobs TBQueue (NonEmpty (JobRead payload))
workQueue TVar Int
busyCount TVar Bool
workerFinishedVar = m () -> m ()
forall (f :: * -> *) a b. Applicative f => f a -> f b
forever (m () -> m ()) -> m () -> m ()
forall a b. (a -> b) -> a -> b
$ m () -> m ()
forall (m :: * -> *) a. MonadUnliftIO m => m a -> m a
mask_ (m () -> m ()) -> m () -> m ()
forall a b. (a -> b) -> a -> b
$ do
  -- Mask covers the window between the atomic claim (which increments
  -- busyCount) and entering the finally block that decrements it.
  jobBatch <- STM (NonEmpty (JobRead payload)) -> m (NonEmpty (JobRead payload))
forall (m :: * -> *) a. MonadIO m => STM a -> m a
atomically (STM (NonEmpty (JobRead payload))
 -> m (NonEmpty (JobRead payload)))
-> STM (NonEmpty (JobRead payload))
-> m (NonEmpty (JobRead payload))
forall a b. (a -> b) -> a -> b
$ do
    batch <- TBQueue (NonEmpty (JobRead payload))
-> STM (NonEmpty (JobRead payload))
forall a. TBQueue a -> STM a
readTBQueue TBQueue (NonEmpty (JobRead payload))
workQueue
    modifyTVar' busyCount (+ 1)
    pure batch

  let jobIds = (JobRead payload -> Int64) -> [JobRead payload] -> [Int64]
forall a b. (a -> b) -> [a] -> [b]
map JobRead payload -> Int64
forall payload key q insertedAt adm.
JobRecord payload key q insertedAt adm -> key
Job.primaryKey (NonEmpty (JobRead payload) -> [JobRead payload]
forall a. NonEmpty a -> [a]
forall (t :: * -> *) a. Foldable t => t a -> [a]
toList NonEmpty (JobRead payload)
jobBatch)

  flip
    finally
    ( atomically $ do
        modifyTVar' busyCount (subtract 1)
        writeTVar workerFinishedVar True
    )
    $ do
      handoff <- newCancelHandoff
      result <-
        withRegisteredJobs runningJobs jobIds $
          processJobsWithRetry config consumeSpan handoff jobBatch
      case result of
        Right () -> () -> m ()
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
        Left SomeException
exception
          -- Finalized inside the job span. A cancel delivered before that catch, or
          -- one that interrupts it, arrives here undone.
          | Just (JobForceCancelled [Int64]
cancelledIds [Int64]
reclaimedIds) <- SomeException -> Maybe JobForceCancelled
forall e. Exception e => SomeException -> Maybe e
fromException SomeException
exception -> do
              alreadyFinalized <- CancelHandoff -> m Bool
forall (m :: * -> *). MonadIO m => CancelHandoff -> m Bool
cancelFinalized CancelHandoff
handoff
              unless alreadyFinalized $
                finalizeForceCancelled config jobBatch cancelledIds reclaimedIds handoff
          | Just AsyncCancelled
Async.AsyncCancelled <- SomeException -> Maybe AsyncCancelled
forall e. Exception e => SomeException -> Maybe e
fromException SomeException
exception -> IO () -> m ()
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (SomeException -> IO ()
forall e a. (HasCallStack, Exception e) => e -> IO a
E.throwIO SomeException
exception)
          | Bool
otherwise -> do
              LogConfig -> LogLevel -> Text -> m ()
forall (m :: * -> *).
MonadUnliftIO m =>
LogConfig -> LogLevel -> Text -> m ()
tryLog (WorkerConfig m payload -> NonEmpty (JobRead payload) -> LogConfig
forall (m :: * -> *) payload.
WorkerConfig m payload -> NonEmpty (JobRead payload) -> LogConfig
batchLog WorkerConfig m payload
config NonEmpty (JobRead payload)
jobBatch) LogLevel
Error (Text -> m ()) -> Text -> m ()
forall a b. (a -> b) -> a -> b
$ Text
"Worker exception: " Text -> Text -> Text
forall a. Semigroup a => a -> a -> a
<> SomeException -> Text
displayEx SomeException
exception
              Int -> m ()
forall (m :: * -> *). MonadIO m => Int -> m ()
threadDelay Int
2_000_000

processJobsWithRetry
  :: forall payload m
   . ( EncodeJobResult (ResultOf m payload)
     , JobOperation m payload
     )
  => WorkerConfig m payload
  -> ConsumeSpan
  -- ^ The pool's consumer-span shape, built once for its queue.
  -> CancelHandoff
  -> NonEmpty (Job.JobRead payload)
  -> m ()
processJobsWithRetry :: forall payload (m :: * -> *).
(EncodeJobResult (ResultOf m payload), JobOperation m payload) =>
WorkerConfig m payload
-> ConsumeSpan
-> CancelHandoff
-> NonEmpty (JobRead payload)
-> m ()
processJobsWithRetry WorkerConfig m payload
config ConsumeSpan
consumeSpan CancelHandoff
handoff NonEmpty (JobRead payload)
jobs = do
  startTime <- IO UTCTime -> m UTCTime
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO IO UTCTime
getCurrentTime
  schemaName <- Arb.getSchema
  tracer <- resolveTracer
  let (firstJob :| _) = jobs
      -- Rethrown with base throwIO. The flag is set last. An interrupted finalizer
      -- leaves the rest to 'workerLoop'.
      onForceCancel exc :: JobForceCancelled
exc@(JobForceCancelled [Int64]
cancelledIds [Int64]
goneIds) = do
        WorkerConfig m payload
-> NonEmpty (JobRead payload)
-> [Int64]
-> [Int64]
-> CancelHandoff
-> m ()
forall (m :: * -> *) payload.
JobOperation m payload =>
WorkerConfig m payload
-> NonEmpty (JobRead payload)
-> [Int64]
-> [Int64]
-> CancelHandoff
-> m ()
finalizeForceCancelled WorkerConfig m payload
config NonEmpty (JobRead payload)
jobs [Int64]
cancelledIds [Int64]
goneIds CancelHandoff
handoff
        CancelHandoff -> m ()
forall (m :: * -> *). MonadIO m => CancelHandoff -> m ()
markCancelFinalized CancelHandoff
handoff
        IO () -> m ()
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO () -> m ()) -> IO () -> m ()
forall a b. (a -> b) -> a -> b
$ JobForceCancelled -> IO ()
forall e a. (HasCallStack, Exception e) => e -> IO a
E.throwIO JobForceCancelled
exc
      claimHook JobRead payload
job =
        LogConfig -> Text -> m () -> m ()
forall (m :: * -> *).
MonadUnliftIO m =>
LogConfig -> Text -> m () -> m ()
runHook (WorkerConfig m payload -> JobRead payload -> LogConfig
forall (m :: * -> *) payload.
WorkerConfig m payload -> JobRead payload -> LogConfig
jobLog WorkerConfig m payload
config JobRead payload
job) Text
"onJobClaimed" (m () -> m ()) -> m () -> m ()
forall a b. (a -> b) -> a -> b
$
          ObservabilityHooks m payload
-> JobPayload payload => JobRead payload -> UTCTime -> m ()
forall (m :: * -> *) payload.
ObservabilityHooks m payload
-> JobPayload payload => JobRead payload -> UTCTime -> m ()
Job.onJobClaimed (WorkerConfig m payload -> ObservabilityHooks m payload
forall (m :: * -> *) payload.
WorkerConfig m payload -> ObservabilityHooks m payload
observabilityHooks WorkerConfig m payload
config) JobRead payload
job UTCTime
startTime
  -- The span covers the claim hooks, the outcome report and the force-cancel
  -- finalizer.
  withConsumeSpan tracer consumeSpan jobs $ flip catchSyncOrAsync onForceCancel $ do
    traverse_ claimHook jobs
    result <-
      tryAny
        $ withJobsHeartbeat
          (observabilityHooks config)
          (jobHeartbeatInterval config)
          (visibilityTimeout config)
          (maxJobDuration config)
          startTime
          jobs
          (pendingJobs handoff jobs)
          (logConfig config)
          (heartbeatSignal config)
        $ case handlerMode config of
          SingleJobMode JobHandler m payload (ResultOf m payload)
handler ->
            CancelHandoff -> Settled payload -> m () -> (() -> m ()) -> m ()
forall (m :: * -> *) payload a b.
MonadIO m =>
CancelHandoff -> Settled payload -> m a -> (a -> m b) -> m b
settleInterruptibly
              CancelHandoff
handoff
              ([JobRead payload] -> Settled payload
forall payload. [JobRead payload] -> Settled payload
finalized [JobRead payload
firstJob])
              ( m () -> m ()
forall a. m a -> m a
forall (m :: * -> *) a. MonadArbiter m => m a -> m a
withDbTransaction (m () -> m ()) -> m () -> m ()
forall a b. (a -> b) -> a -> b
$ do
                  handlerResult <- JobHandler m payload (ResultOf m payload)
-> JobRead payload -> m (ResultOf m payload)
forall payload.
JobHandler m payload (ResultOf m payload)
-> JobRead payload -> m (ResultOf m payload)
forall (m :: * -> *) payload.
MonadArbiter m =>
JobHandler m payload (ResultOf m payload)
-> JobRead payload -> m (ResultOf m payload)
runHandlerWithConnection JobHandler m payload (ResultOf m payload)
handler JobRead payload
firstJob
                  ackOrGone firstJob
                  storeJobResult schemaName firstJob handlerResult
              )
              (m () -> () -> m ()
forall a b. a -> b -> a
const (WorkerConfig m payload -> UTCTime -> JobRead payload -> m ()
forall (m :: * -> *) payload.
JobOperation m payload =>
WorkerConfig m payload -> UTCTime -> JobRead payload -> m ()
reportSuccess WorkerConfig m payload
config UTCTime
startTime JobRead payload
firstJob))
          BatchedJobsMode Int
_ NonEmpty (JobRead payload)
-> BatchCallbacks m payload (ResultOf m payload) -> m ()
handler -> NonEmpty (JobRead payload)
-> BatchCallbacks m payload (ResultOf m payload) -> m ()
handler NonEmpty (JobRead payload)
jobs (WorkerConfig m payload
-> CancelHandoff
-> NonEmpty (JobRead payload)
-> UTCTime
-> Text
-> BatchCallbacks m payload (ResultOf m payload)
forall payload (m :: * -> *).
(EncodeJobResult (ResultOf m payload), JobOperation m payload) =>
WorkerConfig m payload
-> CancelHandoff
-> NonEmpty (JobRead payload)
-> UTCTime
-> Text
-> BatchCallbacks m payload (ResultOf m payload)
batchCallbacks WorkerConfig m payload
config CancelHandoff
handoff NonEmpty (JobRead payload)
jobs UTCTime
startTime Text
schemaName)
    endTime <- liftIO getCurrentTime
    reportBatchOutcome config startTime endTime jobs handoff result