{-# LANGUAGE OverloadedStrings #-}

-- | STM handlers for the worker pool's NOTIFY channels.
-- Each handler decodes the payload and reacts only if the message
-- addresses this worker.
module Arbiter.Worker.ChannelHandlers
  ( RunningJobs
  , handlePauseNotif
  , handleCancelNotif
  , handleCronRunNotif
  , withRegisteredJobs
  ) where

import Arbiter.Core.Exceptions (JobForceCancelled (..))
import Arbiter.Core.Listen (Notification, notificationData)
import Control.Concurrent (forkIO)
import Control.Exception (SomeException)
import Control.Monad (unless, void, when)
import Data.Aeson qualified as Aeson
import Data.Foldable (traverse_)
import Data.Int (Int64)
import Data.Map.Strict qualified as Map
import Data.Set (Set)
import Data.Set qualified as Set
import Data.Text (Text)
import Data.Text.Encoding (decodeUtf8Lenient)
import Data.UUID (UUID)
import UnliftIO (MonadUnliftIO, atomically, liftIO)
import UnliftIO.Async qualified as Async
import UnliftIO.Exception (finally, throwTo)
import UnliftIO.STM (TVar, readTVar)
import UnliftIO.STM qualified as STM

import Arbiter.Worker.Config (WorkerConfig (..), workerStateVar, writePause)
import Arbiter.Worker.WorkerState (WorkerState (..))

-- | The handler threads in flight, by job id.
type RunningJobs = TVar (Map.Map Int64 (Async.Async ()))

-- | Decode the pause payload and, if it addresses this worker, write 'pauseVar'.
handlePauseNotif
  :: (MonadUnliftIO m)
  => WorkerConfig n payload
  -> Notification
  -> m ()
handlePauseNotif :: forall (m :: * -> *) (n :: * -> *) payload.
MonadUnliftIO m =>
WorkerConfig n payload -> Notification -> m ()
handlePauseNotif WorkerConfig n payload
config Notification
notif =
  case ByteString -> Maybe PausePayload
forall a. FromJSON a => ByteString -> Maybe a
Aeson.decodeStrict (Notification -> ByteString
notificationData Notification
notif) :: Maybe PausePayload of
    Just (PausePayload UUID
wid Bool
paused) | UUID
wid UUID -> UUID -> Bool
forall a. Eq a => a -> a -> Bool
== WorkerConfig n payload -> UUID
forall (m :: * -> *) payload. WorkerConfig m payload -> UUID
workerId WorkerConfig n payload
config -> STM () -> m ()
forall (m :: * -> *) a. MonadIO m => STM a -> m a
atomically (STM () -> m ()) -> STM () -> m ()
forall a b. (a -> b) -> a -> b
$ do
      state <- TVar WorkerState -> STM WorkerState
forall a. TVar a -> STM a
STM.readTVar (WorkerConfig n payload -> TVar WorkerState
forall (n :: * -> *) payload.
WorkerConfig n payload -> TVar WorkerState
workerStateVar WorkerConfig n payload
config)
      unless (state == ShuttingDown) $ writePause config paused
    Maybe PausePayload
_ -> () -> m ()
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()

-- | If the cancel payload targets a job on this worker, 'throwTo'
-- 'JobForceCancelled' into its handler thread.
handleCancelNotif
  :: (MonadUnliftIO m)
  => WorkerConfig n payload
  -> RunningJobs
  -> Notification
  -> m ()
handleCancelNotif :: forall (m :: * -> *) (n :: * -> *) payload.
MonadUnliftIO m =>
WorkerConfig n payload -> RunningJobs -> Notification -> m ()
handleCancelNotif WorkerConfig n payload
config RunningJobs
runningJobs Notification
notif =
  case ByteString -> Maybe CancelPayload
forall a. FromJSON a => ByteString -> Maybe a
Aeson.decodeStrict (Notification -> ByteString
notificationData Notification
notif) :: Maybe CancelPayload of
    Just (CancelPayload UUID
wid Int64
jid) | UUID
wid UUID -> UUID -> Bool
forall a. Eq a => a -> a -> Bool
== WorkerConfig n payload -> UUID
forall (m :: * -> *) payload. WorkerConfig m payload -> UUID
workerId WorkerConfig n payload
config -> do
      mAsync <- STM (Maybe (Async ())) -> m (Maybe (Async ()))
forall (m :: * -> *) a. MonadIO m => STM a -> m a
atomically (STM (Maybe (Async ())) -> m (Maybe (Async ())))
-> STM (Maybe (Async ())) -> m (Maybe (Async ()))
forall a b. (a -> b) -> a -> b
$ Int64 -> Map Int64 (Async ()) -> Maybe (Async ())
forall k a. Ord k => k -> Map k a -> Maybe a
Map.lookup Int64
jid (Map Int64 (Async ()) -> Maybe (Async ()))
-> STM (Map Int64 (Async ())) -> STM (Maybe (Async ()))
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> RunningJobs -> STM (Map Int64 (Async ()))
forall a. TVar a -> STM a
readTVar RunningJobs
runningJobs
      traverse_ (fireCancel jid) mAsync
    Maybe CancelPayload
_ -> () -> m ()
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
  where
    -- Fork the throwTo off the listener thread.
    fireCancel :: Int64 -> Async a -> m ()
fireCancel Int64
jid Async a
handlerAsync =
      IO () -> m ()
forall a. IO a -> m a
forall (m :: * -> *) a. MonadIO m => IO a -> m a
liftIO (IO () -> m ()) -> (IO () -> IO ()) -> IO () -> m ()
forall b c a. (b -> c) -> (a -> b) -> a -> c
. IO ThreadId -> IO ()
forall (f :: * -> *) a. Functor f => f a -> f ()
void (IO ThreadId -> IO ()) -> (IO () -> IO ThreadId) -> IO () -> IO ()
forall b c a. (b -> c) -> (a -> b) -> a -> c
. IO () -> IO ThreadId
forkIO (IO () -> m ()) -> IO () -> m ()
forall a b. (a -> b) -> a -> b
$ ThreadId -> JobForceCancelled -> IO ()
forall e (m :: * -> *).
(Exception e, MonadIO m) =>
ThreadId -> e -> m ()
throwTo (Async a -> ThreadId
forall a. Async a -> ThreadId
Async.asyncThreadId Async a
handlerAsync) ([Int64] -> [Int64] -> JobForceCancelled
JobForceCancelled [Int64
jid] [])

-- | Signal the scheduler when a run-now NOTIFY names a schedule this pool owns.
handleCronRunNotif
  :: (MonadUnliftIO m)
  => Set Text
  -- ^ This pool's own cron schedule names
  -> TVar Bool
  -> Notification
  -> m ()
handleCronRunNotif :: forall (m :: * -> *).
MonadUnliftIO m =>
Set Text -> TVar Bool -> Notification -> m ()
handleCronRunNotif Set Text
ownNames TVar Bool
runNowVar Notification
notif =
  Bool -> m () -> m ()
forall (f :: * -> *). Applicative f => Bool -> f () -> f ()
when (Text -> Set Text -> Bool
forall a. Ord a => a -> Set a -> Bool
Set.member (ByteString -> Text
decodeUtf8Lenient (Notification -> ByteString
notificationData Notification
notif)) Set Text
ownNames)
    (m () -> m ()) -> m () -> m ()
forall a b. (a -> b) -> a -> b
$ STM () -> m ()
forall (m :: * -> *) a. MonadIO m => STM a -> m a
atomically
    (STM () -> m ()) -> STM () -> m ()
forall a b. (a -> b) -> a -> b
$ TVar Bool -> Bool -> STM ()
forall a. TVar a -> a -> STM ()
STM.writeTVar TVar Bool
runNowVar Bool
True

-- | Run @work@ as an async registered in 'RunningJobs' for its lifetime. Cleanup
-- removes the entries this call made.
withRegisteredJobs
  :: forall m
   . (MonadUnliftIO m)
  => RunningJobs
  -> [Int64]
  -> m ()
  -> m (Either SomeException ())
withRegisteredJobs :: forall (m :: * -> *).
MonadUnliftIO m =>
RunningJobs -> [Int64] -> m () -> m (Either SomeException ())
withRegisteredJobs RunningJobs
runningJobs [Int64]
jobIds m ()
work = do
  startGate <- m (TMVar ())
forall (m :: * -> *) a. MonadIO m => m (TMVar a)
STM.newEmptyTMVarIO
  let gated :: (forall c. m c -> m c) -> m ()
      gated forall c. m c -> m c
unmask = m () -> m ()
forall c. m c -> m c
unmask (m () -> m ()) -> m () -> m ()
forall a b. (a -> b) -> a -> b
$ do
        STM () -> m ()
forall (m :: * -> *) a. MonadIO m => STM a -> m a
atomically (TMVar () -> STM ()
forall a. TMVar a -> STM a
STM.readTMVar TMVar ()
startGate)
        m ()
work
      unregister Async ()
handlerAsync =
        STM () -> m ()
forall (m :: * -> *) a. MonadIO m => STM a -> m a
atomically (STM () -> m ()) -> STM () -> m ()
forall a b. (a -> b) -> a -> b
$
          RunningJobs
-> (Map Int64 (Async ()) -> Map Int64 (Async ())) -> STM ()
forall a. TVar a -> (a -> a) -> STM ()
STM.modifyTVar' RunningJobs
runningJobs ((Map Int64 (Async ()) -> Map Int64 (Async ())) -> STM ())
-> (Map Int64 (Async ()) -> Map Int64 (Async ())) -> STM ()
forall a b. (a -> b) -> a -> b
$ \Map Int64 (Async ())
running ->
            (Map Int64 (Async ()) -> Int64 -> Map Int64 (Async ()))
-> Map Int64 (Async ()) -> [Int64] -> Map Int64 (Async ())
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl'
              (\Map Int64 (Async ())
acc Int64
jid -> (Async () -> Maybe (Async ()))
-> Int64 -> Map Int64 (Async ()) -> Map Int64 (Async ())
forall k a. Ord k => (a -> Maybe a) -> k -> Map k a -> Map k a
Map.update (\Async ()
registered -> if Async ()
registered Async () -> Async () -> Bool
forall a. Eq a => a -> a -> Bool
== Async ()
handlerAsync then Maybe (Async ())
forall a. Maybe a
Nothing else Async () -> Maybe (Async ())
forall a. a -> Maybe a
Just Async ()
registered) Int64
jid Map Int64 (Async ())
acc)
              Map Int64 (Async ())
running
              [Int64]
jobIds
  Async.withAsyncWithUnmask gated $ \Async ()
handlerAsync ->
    (m (Either SomeException ())
 -> m () -> m (Either SomeException ()))
-> m ()
-> m (Either SomeException ())
-> m (Either SomeException ())
forall a b c. (a -> b -> c) -> b -> a -> c
flip m (Either SomeException ()) -> m () -> m (Either SomeException ())
forall (m :: * -> *) a b. MonadUnliftIO m => m a -> m b -> m a
finally (Async () -> m ()
unregister Async ()
handlerAsync) (m (Either SomeException ()) -> m (Either SomeException ()))
-> m (Either SomeException ()) -> m (Either SomeException ())
forall a b. (a -> b) -> a -> b
$ do
      STM () -> m ()
forall (m :: * -> *) a. MonadIO m => STM a -> m a
atomically (STM () -> m ()) -> STM () -> m ()
forall a b. (a -> b) -> a -> b
$ do
        RunningJobs
-> (Map Int64 (Async ()) -> Map Int64 (Async ())) -> STM ()
forall a. TVar a -> (a -> a) -> STM ()
STM.modifyTVar' RunningJobs
runningJobs ((Map Int64 (Async ()) -> Map Int64 (Async ())) -> STM ())
-> (Map Int64 (Async ()) -> Map Int64 (Async ())) -> STM ()
forall a b. (a -> b) -> a -> b
$ \Map Int64 (Async ())
running ->
          (Map Int64 (Async ()) -> Int64 -> Map Int64 (Async ()))
-> Map Int64 (Async ()) -> [Int64] -> Map Int64 (Async ())
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' (\Map Int64 (Async ())
acc Int64
jid -> Int64 -> Async () -> Map Int64 (Async ()) -> Map Int64 (Async ())
forall k a. Ord k => k -> a -> Map k a -> Map k a
Map.insert Int64
jid Async ()
handlerAsync Map Int64 (Async ())
acc) Map Int64 (Async ())
running [Int64]
jobIds
        TMVar () -> () -> STM ()
forall a. TMVar a -> a -> STM ()
STM.putTMVar TMVar ()
startGate ()
      Async () -> m (Either SomeException ())
forall (m :: * -> *) a.
MonadIO m =>
Async a -> m (Either SomeException a)
Async.waitCatch Async ()
handlerAsync

data PausePayload = PausePayload UUID Bool

instance Aeson.FromJSON PausePayload where
  parseJSON :: Value -> Parser PausePayload
parseJSON = String
-> (Object -> Parser PausePayload) -> Value -> Parser PausePayload
forall a. String -> (Object -> Parser a) -> Value -> Parser a
Aeson.withObject String
"PausePayload" ((Object -> Parser PausePayload) -> Value -> Parser PausePayload)
-> (Object -> Parser PausePayload) -> Value -> Parser PausePayload
forall a b. (a -> b) -> a -> b
$ \Object
obj ->
    UUID -> Bool -> PausePayload
PausePayload (UUID -> Bool -> PausePayload)
-> Parser UUID -> Parser (Bool -> PausePayload)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Object
obj Object -> Key -> Parser UUID
forall a. FromJSON a => Object -> Key -> Parser a
Aeson..: Key
"worker_id" Parser (Bool -> PausePayload) -> Parser Bool -> Parser PausePayload
forall a b. Parser (a -> b) -> Parser a -> Parser b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> Object
obj Object -> Key -> Parser Bool
forall a. FromJSON a => Object -> Key -> Parser a
Aeson..: Key
"paused"

data CancelPayload = CancelPayload UUID Int64

instance Aeson.FromJSON CancelPayload where
  parseJSON :: Value -> Parser CancelPayload
parseJSON = String
-> (Object -> Parser CancelPayload)
-> Value
-> Parser CancelPayload
forall a. String -> (Object -> Parser a) -> Value -> Parser a
Aeson.withObject String
"CancelPayload" ((Object -> Parser CancelPayload) -> Value -> Parser CancelPayload)
-> (Object -> Parser CancelPayload)
-> Value
-> Parser CancelPayload
forall a b. (a -> b) -> a -> b
$ \Object
obj ->
    UUID -> Int64 -> CancelPayload
CancelPayload (UUID -> Int64 -> CancelPayload)
-> Parser UUID -> Parser (Int64 -> CancelPayload)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Object
obj Object -> Key -> Parser UUID
forall a. FromJSON a => Object -> Key -> Parser a
Aeson..: Key
"worker_id" Parser (Int64 -> CancelPayload)
-> Parser Int64 -> Parser CancelPayload
forall a b. Parser (a -> b) -> Parser a -> Parser b
forall (f :: * -> *) a b. Applicative f => f (a -> b) -> f a -> f b
<*> Object
obj Object -> Key -> Parser Int64
forall a. FromJSON a => Object -> Key -> Parser a
Aeson..: Key
"job_id"