packages feed

marionette-1.1.0: src/Test/Marionette/Client.hs

{-# LANGUAGE RoleAnnotations #-}
{-# LANGUAGE UndecidableInstances #-}
{-# OPTIONS_GHC -Wno-orphans #-}

module Test.Marionette.Client where

import Control.Applicative ((<|>))
import Control.Exception (AssertionFailed (AssertionFailed), SomeException)
import Control.Monad (void, (<=<))
import Control.Monad.Catch
    ( Exception (..)
    , MonadCatch
    , MonadMask
    , MonadThrow
    , catch
    , catchAll
    , finally
    , throwM
    )
import Control.Monad.Error.Class (MonadError (..))
import Control.Monad.IO.Class (MonadIO (liftIO))
import Control.Monad.Reader (ReaderT (runReaderT))
import Control.Monad.Reader qualified as Reader
import Control.Monad.Reader.Class (MonadReader)
import Data.Aeson (AesonException (AesonException), FromJSON)
import Data.Aeson qualified as Aeson
import Data.Aeson.Types qualified as Aeson
import Data.Binary (Binary)
import Data.Binary qualified as Binary
import Data.Binary.Get qualified as Binary
import Data.ByteString (ByteString)
import Data.ByteString qualified as BS
import Data.ByteString.Builder.Extra qualified as ByteString
import Data.IntMap (IntMap)
import Data.IntMap.Strict qualified as IntMap
import GHC.Stack (HasCallStack)
import Network.Simple.TCP
    ( HostName
    , ServiceName
    , SockAddr
    , Socket
    , closeSock
    , connectSock
    , sendLazy
    )
import Network.Socket (PortNumber)
import Network.Socket.ByteString (recv)
import System.Timeout (timeout)
import Test.Marionette.Class (Marionette (..))
import Test.Marionette.Protocol
import UnliftIO
    ( MonadUnliftIO
    , TMVar
    , TQueue
    , async
    , bracket
    , isEmptyTMVar
    , link
    , readTMVar
    )
import UnliftIO.Retry (constantDelay, limitRetriesByCumulativeDelay, recoverAll)
import UnliftIO.STM
    ( TVar
    , atomically
    , modifyTVar'
    , newEmptyTMVarIO
    , newTQueueIO
    , newTVarIO
    , putTMVar
    , readTQueue
    , stateTVar
    , tryPutTMVar
    , writeTQueue
    )
import Prelude hiding (log)

data SocketClosed = SocketClosed
    deriving stock (Show)
    deriving anyclass (Exception)

newtype DecodeError = DecodeError String
    deriving stock (Show)
    deriving anyclass (Exception)

incoming
    :: forall m a
     . (MonadUnliftIO m, MonadCatch m, Binary a)
    => Socket
    -> (SomeException -> m ())
    -> m (TQueue a)
incoming socket onFailure = do
    q <- newTQueueIO
    void . async $ go (atomically . writeTQueue q) (newDecoder "") `catch` onFailure
    pure q
  where
    newDecoder :: ByteString -> Binary.Decoder a
    newDecoder bs = Binary.runGetIncremental Binary.get `Binary.pushChunk` bs
    go _ (Binary.Fail _ _ err) = throwM . DecodeError $ err
    go f (Binary.Done rest _ a) = do
        f a
        go f $ newDecoder rest
    go f dec = do
        chunk <-
            liftIO (recv socket ByteString.defaultChunkSize `catchAll` \_ -> throwM SocketClosed)
        if BS.null chunk
            then throwM SocketClosed
            else go f (dec `Binary.pushChunk` chunk)

connect
    :: (MonadUnliftIO m)
    => HostName
    -> ServiceName
    -> ((Socket, SockAddr) -> m a)
    -> m a
connect host port =
    bracket
        ( recoverAll (limitRetriesByCumulativeDelay 5_000_000 $ constantDelay 50_000)
            . const
            $ connectSock host port
        )
        (closeSock . fst)

newtype MarionetteTimeout = MarionetteTimeout Command
    deriving stock (Show)
    deriving anyclass (Exception)

data QueuedCommand = QueuedCommand Int Command

data ClientEnv = ClientEnv
    { sendQueue :: TQueue QueuedCommand
    , connectionLost :: TMVar SomeException
    , commandTimeout :: Int
    , nextMessageId :: TVar Int
    , pendingCommands :: TVar (IntMap (TMVar Result))
    }

defaultCommandTimeout :: Int
defaultCommandTimeout = 15_000_000

-- | A monad transformer that speaks the Marionette wire protocol over a TCP socket.
-- Run it with 'runMarionetteT'.
newtype MarionetteT m a = MarionetteT (ReaderT ClientEnv m a)
    deriving newtype
        ( Functor
        , Applicative
        , Monad
        , MonadThrow
        , MonadCatch
        , MonadMask
        , MonadIO
        , MonadUnliftIO
        , MonadReader ClientEnv
        )

instance (MonadUnliftIO m, MonadThrow m, MonadCatch m) => Marionette (MarionetteT m) where
    sendCommand :: (HasCallStack, FromJSON a) => Command -> MarionetteT m a
    sendCommand command = do
        ClientEnv{..} <- Reader.ask
        result <- newEmptyTMVarIO
        messageId <- atomically do
            messageId <- stateTVar nextMessageId $ \n -> (n, n + 1)
            modifyTVar' pendingCommands $ IntMap.insert messageId result
            writeTQueue sendQueue $ QueuedCommand messageId command
            pure messageId
        outcome <-
            liftIO . timeout commandTimeout . atomically $
                (Right <$> readTMVar result) <|> (Left <$> readTMVar connectionLost)
        case outcome of
            Nothing -> atomically . modifyTVar' pendingCommands $ IntMap.delete messageId
            Just _ -> pure ()
        maybe
            (throwM . MarionetteTimeout $ command)
            (either throwM (either throwM pure <=< parseResult))
            outcome
      where
        parseResult :: (FromJSON a) => Result -> MarionetteT m (Either Error a)
        parseResult =
            either (pure . Left) $
                either (throwM . AesonException) (pure . Right)
                    . Aeson.parseEither Aeson.parseJSON

instance (MonadThrow m, MonadCatch m) => MonadError Error (MarionetteT m) where
    throwError = throwM
    catchError = catch

-- | Run an action against a Marionette server.
runMarionetteTWith
    :: forall m a
     . (MonadUnliftIO m, MonadMask m)
    => HostName
    -> PortNumber
    -> Int
    -- ^ Command timeout, in microseconds.
    -> MarionetteT m a
    -> m a
runMarionetteTWith host port commandTimeout action = do
    sendQueue :: TQueue QueuedCommand <- newTQueueIO
    connectionLost <- newEmptyTMVarIO
    nextMessageId <- newTVarIO 1
    pendingCommands <- newTVarIO mempty
    done <- newEmptyTMVarIO
    let repeatUntilDone :: forall m'. (MonadIO m') => m' () -> m' ()
        repeatUntilDone a =
            atomically (isEmptyTMVar done) >>= \case
                False -> pure ()
                True -> a >> repeatUntilDone a
        handleIncoming :: MarionetteMessage -> MarionetteT m ()
        handleIncoming message =
            decodeMarionetteM message >>= \Message{..} ->
                atomically
                    ( stateTVar pendingCommands $
                        IntMap.updateLookupWithKey (\_ _ -> Nothing) messageId
                    )
                    >>= \case
                        Nothing -> pure () -- a late or unsolicited reply; not fatal
                        Just result -> atomically $ putTMVar result messageContent
        failPending :: SomeException -> MarionetteT m ()
        failPending = void . atomically . tryPutTMVar connectionLost
        handleCommand :: Socket -> QueuedCommand -> MarionetteT m ()
        handleCommand socket (QueuedCommand messageId messageContent) =
            liftIO . sendLazy socket . Binary.encode . MarionetteMessage . Aeson.encode $
                Message{..}
        runSocket :: MarionetteT m ()
        runSocket =
            connect host (show port) \(socket, _) -> do
                incomingQueue <- incoming socket failPending
                void . decodeMarionetteM @_ @Greeting =<< atomically (readTQueue incomingQueue)
                let send =
                        atomically
                            ( isEmptyTMVar done >>= \case
                                True -> Just <$> readTQueue sendQueue
                                False -> pure Nothing
                            )
                            >>= maybe (pure ()) ((>> send) . handleCommand socket)
                linkedAsync send
                void . async $
                    repeatUntilDone (handleIncoming =<< atomically (readTQueue incomingQueue))
                        `catch` failPending
                atomically $ readTMVar done
    flip runReaderT ClientEnv{..} . (\(MarionetteT r) -> r) $ do
        linkedAsync runSocket
        action `finally` atomically (void $ tryPutTMVar done ())
  where
    linkedAsync :: MarionetteT m () -> MarionetteT m ()
    linkedAsync = link <=< async

    decodeMarionetteM :: forall m' a'. (MonadThrow m', FromJSON a') => MarionetteMessage -> m' a'
    decodeMarionetteM = either (throwM . AssertionFailed) pure . decodeMarionette

    decodeMarionette :: forall a'. (FromJSON a') => MarionetteMessage -> Either String a'
    decodeMarionette (MarionetteMessage lbs) = Aeson.eitherDecode lbs

-- | Run an action against a Marionette server listening on @localhost:2828@
-- (started by launching Firefox with the @--marionette@ flag).
runMarionetteT :: (MonadUnliftIO m, MonadMask m) => MarionetteT m a -> m a
runMarionetteT = runMarionetteTWith "localhost" 2828 defaultCommandTimeout