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