From f76a5ca5b6be8bf0ba1323c3e6b90b17737dba76 Mon Sep 17 00:00:00 2001 From: Evgeny Poberezkin <2769109+epoberezkin@users.noreply.github.com> Date: Sun, 9 Jul 2023 18:04:45 +0100 Subject: [PATCH] agent: catch IO errors correctly in MonadError (#795) * agent: catch IO errors correctly in MonadError * correction * correction * utils * agentFinally to catch IO exceptions in ExceptT * rename * remove, inline * rename utils * utils unit test * test to show catch and finally problems * tryAllErrors * enable all tests --- simplexmq.cabal | 1 + src/Simplex/FileTransfer/Agent.hs | 29 +++-- src/Simplex/Messaging/Agent.hs | 59 +++++----- src/Simplex/Messaging/Agent/Client.hs | 4 +- src/Simplex/Messaging/Agent/Env/SQLite.hs | 23 +++- .../Messaging/Agent/NtfSubSupervisor.hs | 10 +- src/Simplex/Messaging/Util.hs | 31 +++-- tests/CoreTests/UtilTests.hs | 110 ++++++++++++++++++ tests/Test.hs | 2 + 9 files changed, 209 insertions(+), 60 deletions(-) create mode 100644 tests/CoreTests/UtilTests.hs diff --git a/simplexmq.cabal b/simplexmq.cabal index c9506f994..3c20ea9ed 100644 --- a/simplexmq.cabal +++ b/simplexmq.cabal @@ -533,6 +533,7 @@ test-suite simplexmq-test CoreTests.EncodingTests CoreTests.ProtocolErrorTests CoreTests.RetryIntervalTests + CoreTests.UtilTests CoreTests.VersionRangeTests FileDescriptionTests NtfClient diff --git a/src/Simplex/FileTransfer/Agent.hs b/src/Simplex/FileTransfer/Agent.hs index b50031f96..7ca1e5a4f 100644 --- a/src/Simplex/FileTransfer/Agent.hs +++ b/src/Simplex/FileTransfer/Agent.hs @@ -71,7 +71,6 @@ import System.FilePath (takeFileName, ()) import UnliftIO import UnliftIO.Concurrent import UnliftIO.Directory -import qualified UnliftIO.Exception as E startWorkers :: AgentMonad m => AgentClient -> Maybe FilePath -> m () startWorkers c workDir = do @@ -162,7 +161,7 @@ addWorker c wsSel runWorker runWorkerNoSrv srv_ = do let runWorker' = case srv_ of Just srv -> runWorker c srv doWork Nothing -> runWorkerNoSrv c doWork - worker <- async $ runWorker' `E.finally` atomically (TM.delete srv_ ws) + worker <- async $ runWorker' `agentFinally` atomically (TM.delete srv_ ws) atomically $ TM.insert srv_ (doWork, worker) ws Just (doWork, _) -> void . atomically $ tryPutTMVar doWork () @@ -187,10 +186,10 @@ runXFTPRcvWorker c srv doWork = do let ri' = maybe ri (\d -> ri {initialInterval = d, increaseAfter = 0}) delay withRetryInterval ri' $ \delay' loop -> downloadFileChunk fc replica - `catchError` \e -> retryOnError "XFTP rcv worker" (retryLoop loop e delay') (retryDone e) e + `catchAgentError` \e -> retryOnError "XFTP rcv worker" (retryLoop loop e delay') (retryDone e) e where retryLoop loop e replicaDelay = do - flip catchError (\_ -> pure ()) $ do + flip catchAgentError (\_ -> pure ()) $ do notifyOnRetry <- asks (xftpNotifyErrsOnRetry . config) when notifyOnRetry $ notify c rcvFileEntityId $ RFERR e closeXFTPServerClient c userId server digest @@ -249,7 +248,7 @@ runXFTPRcvLocalWorker c doWork = do case nextFile of Nothing -> noWorkToDo Just f@RcvFile {rcvFileId, rcvFileEntityId, tmpPath} -> - decryptFile f `catchError` (rcvWorkerInternalError c rcvFileId rcvFileEntityId tmpPath . show) + decryptFile f `catchAgentError` (rcvWorkerInternalError c rcvFileId rcvFileEntityId tmpPath . show) noWorkToDo = void . atomically $ tryTakeTMVar doWork decryptFile :: RcvFile -> m () decryptFile RcvFile {rcvFileId, rcvFileEntityId, key, nonce, tmpPath, savePath, status, chunks} = do @@ -300,7 +299,7 @@ sendFileExperimental c@AgentClient {xftpServers} userId filePath numRecipients = createDirectory outputDir let tempPath = workPath "snd" createDirectoryIfMissing False tempPath - runSend fileName outputDir tempPath `catchError` \e -> do + runSend fileName outputDir tempPath `catchAgentError` \e -> do cleanup outputDir tempPath notify c sndFileId $ SFERR e where @@ -370,7 +369,7 @@ runXFTPSndPrepareWorker c doWork = do case nextFile of Nothing -> noWorkToDo Just f@SndFile {sndFileId, sndFileEntityId, prefixPath} -> - prepareFile f `catchError` (sndWorkerInternalError c sndFileId sndFileEntityId prefixPath . show) + prepareFile f `catchAgentError` (sndWorkerInternalError c sndFileId sndFileEntityId prefixPath . show) noWorkToDo = void . atomically $ tryTakeTMVar doWork prepareFile :: SndFile -> m () prepareFile SndFile {prefixPath = Nothing} = @@ -424,7 +423,7 @@ runXFTPSndPrepareWorker c doWork = do usedSrvs <- newTVarIO ([] :: [XFTPServer]) withRetryInterval (riFast ri) $ \_ loop -> createWithNextSrv usedSrvs - `catchError` \e -> retryOnError "XFTP prepare worker" (retryLoop loop) (throwError e) e + `catchAgentError` \e -> retryOnError "XFTP prepare worker" (retryLoop loop) (throwError e) e where retryLoop loop = atomically (assertAgentForeground c) >> loop createWithNextSrv usedSrvs = do @@ -460,10 +459,10 @@ runXFTPSndWorker c srv doWork = do let ri' = maybe ri (\d -> ri {initialInterval = d, increaseAfter = 0}) delay withRetryInterval ri' $ \delay' loop -> uploadFileChunk fc replica - `catchError` \e -> retryOnError "XFTP snd worker" (retryLoop loop e delay') (retryDone e) e + `catchAgentError` \e -> retryOnError "XFTP snd worker" (retryLoop loop e delay') (retryDone e) e where retryLoop loop e replicaDelay = do - flip catchError (\_ -> pure ()) $ do + flip catchAgentError (\_ -> pure ()) $ do notifyOnRetry <- asks (xftpNotifyErrsOnRetry . config) when notifyOnRetry $ notify c sndFileEntityId $ SFERR e closeXFTPServerClient c userId server digest @@ -579,8 +578,8 @@ deleteSndFileInternal c sndFileEntityId = do deleteSndFileRemote :: forall m. AgentMonad m => AgentClient -> UserId -> SndFileId -> ValidFileDescription 'FSender -> m () deleteSndFileRemote c userId sndFileEntityId (ValidFileDescription FileDescription {chunks}) = do - deleteSndFileInternal c sndFileEntityId `catchError` (notify c sndFileEntityId . SFERR) - forM_ chunks $ \ch -> deleteFileChunk ch `catchError` (notify c sndFileEntityId . SFERR) + deleteSndFileInternal c sndFileEntityId `catchAgentError` (notify c sndFileEntityId . SFERR) + forM_ chunks $ \ch -> deleteFileChunk ch `catchAgentError` (notify c sndFileEntityId . SFERR) where deleteFileChunk :: FileChunk -> m () deleteFileChunk FileChunk {digest, replicas = replica@FileChunkReplica {server} : _} = do @@ -594,7 +593,7 @@ addXFTPDelWorker c srv = do atomically (TM.lookup srv ws) >>= \case Nothing -> do doWork <- newTMVarIO () - worker <- async $ runXFTPDelWorker c srv doWork `E.finally` atomically (TM.delete srv ws) + worker <- async $ runXFTPDelWorker c srv doWork `agentFinally` atomically (TM.delete srv ws) atomically $ TM.insert srv (doWork, worker) ws Just (doWork, _) -> void . atomically $ tryPutTMVar doWork () @@ -619,10 +618,10 @@ runXFTPDelWorker c srv doWork = do let ri' = maybe ri (\d -> ri {initialInterval = d, increaseAfter = 0}) delay withRetryInterval ri' $ \delay' loop -> deleteChunkReplica replica - `catchError` \e -> retryOnError "XFTP del worker" (retryLoop loop e delay') (retryDone e) e + `catchAgentError` \e -> retryOnError "XFTP del worker" (retryLoop loop e delay') (retryDone e) e where retryLoop loop e replicaDelay = do - flip catchError (\_ -> pure ()) $ do + flip catchAgentError (\_ -> pure ()) $ do notifyOnRetry <- asks (xftpNotifyErrsOnRetry . config) when notifyOnRetry $ notify c "" $ SFERR e closeXFTPServerClient c userId server chunkDigest diff --git a/src/Simplex/Messaging/Agent.hs b/src/Simplex/Messaging/Agent.hs index 7c0419e0a..ab09a3a14 100644 --- a/src/Simplex/Messaging/Agent.hs +++ b/src/Simplex/Messaging/Agent.hs @@ -502,7 +502,7 @@ acceptContactAsync' c corrId enableNtfs invId ownConnInfo = do withStore c (`getConn` contactConnId) >>= \case SomeConn _ (ContactConnection ConnData {userId} _) -> do withStore' c $ \db -> acceptInvitation db invId ownConnInfo - joinConnAsync c userId corrId enableNtfs connReq ownConnInfo `catchError` \err -> do + joinConnAsync c userId corrId enableNtfs connReq ownConnInfo `catchAgentError` \err -> do withStore' c (`unacceptInvitation` invId) throwError err _ -> throwError $ CMD PROHIBITED @@ -565,7 +565,7 @@ newConnSrv c userId connId enableNtfs cMode clientData srv = do newRcvConnSrv :: AgentMonad m => AgentClient -> UserId -> ConnId -> Bool -> SConnectionMode c -> Maybe CRClientData -> SMPServerWithAuth -> m (ConnId, ConnectionRequestUri c) newRcvConnSrv c userId connId enableNtfs cMode clientData srv = do AgentConfig {smpClientVRange, smpAgentVRange, e2eEncryptVRange} <- asks config - (rq, qUri) <- newRcvQueue c userId connId srv smpClientVRange `catchError` \e -> liftIO (print e) >> throwError e + (rq, qUri) <- newRcvQueue c userId connId srv smpClientVRange `catchAgentError` \e -> liftIO (print e) >> throwError e void . withStore c $ \db -> updateNewConnRcv db connId rq addSubscription c rq when enableNtfs $ do @@ -671,7 +671,7 @@ acceptContact' c connId enableNtfs invId ownConnInfo = withConnLock c connId "ac withStore c (`getConn` contactConnId) >>= \case SomeConn _ (ContactConnection ConnData {userId} _) -> do withStore' c $ \db -> acceptInvitation db invId ownConnInfo - joinConn c userId connId False enableNtfs connReq ownConnInfo `catchError` \err -> do + joinConn c userId connId False enableNtfs connReq ownConnInfo `catchAgentError` \err -> do withStore' c (`unacceptInvitation` invId) throwError err _ -> throwError $ CMD PROHIBITED @@ -787,7 +787,7 @@ getNotificationMessage' c nonce encNtfInfo = do ntfData <- agentCbDecrypt dhSecret nonce encNtfInfo PNMessageData {smpQueue, ntfTs, nmsgNonce, encNMsgMeta} <- liftEither (parse strP (INTERNAL "error parsing PNMessageData") ntfData) (ntfConnId, rcvNtfDhSecret) <- withStore c (`getNtfRcvQueue` smpQueue) - ntfMsgMeta <- (eitherToMaybe . smpDecode <$> agentCbDecrypt rcvNtfDhSecret nmsgNonce encNMsgMeta) `catchError` \_ -> pure Nothing + ntfMsgMeta <- (eitherToMaybe . smpDecode <$> agentCbDecrypt rcvNtfDhSecret nmsgNonce encNMsgMeta) `catchAgentError` \_ -> pure Nothing maxMsgs <- asks $ ntfMaxMessages . config (NotificationInfo {ntfConnId, ntfTs, ntfMsgMeta},) <$> getNtfMessages ntfConnId maxMsgs ntfMsgMeta [] _ -> throwError $ CMD PROHIBITED @@ -872,8 +872,8 @@ runCommandProcessing c@AgentClient {subQ} server_ = do atomically $ throwWhenInactive c cmdId <- atomically $ readTQueue cq atomically $ beginAgentOperation c AOSndNetwork - E.try (withStore c $ \db -> getPendingCommand db cmdId) >>= \case - Left (e :: E.SomeException) -> atomically $ writeTBQueue subQ ("", "", APC SAEConn $ ERR $ INTERNAL $ show e) + tryAgentError (withStore c $ \db -> getPendingCommand db cmdId) >>= \case + Left e -> atomically $ writeTBQueue subQ ("", "", APC SAEConn $ ERR e) Right cmd -> processCmd (riFast ri) cmdId cmd where processCmd :: RetryInterval -> AsyncCmdId -> PendingCommand -> m () @@ -1078,9 +1078,8 @@ runSmpQueueMsgDelivery c@AgentClient {subQ} cData@ConnData {userId, connId, dupl atomically $ beginAgentOperation c AOSndNetwork atomically $ endAgentOperation c AOMsgDelivery -- this operation begins in queuePendingMsgs let mId = unId msgId - E.try (withStore c $ \db -> getPendingMsgData db connId msgId) >>= \case - Left (e :: E.SomeException) -> - notify $ MERR mId (INTERNAL $ show e) + tryAgentError (withStore c $ \db -> getPendingMsgData db connId msgId) >>= \case + Left e -> notify $ MERR mId e Right (rq_, PendingMsgData {msgType, msgBody, msgFlags, msgRetryState, internalTs}) -> do let ri' = maybe id updateRetryInterval2 msgRetryState ri withRetryLock2 ri' qLock $ \riState loop -> do @@ -1310,7 +1309,7 @@ synchronizeRatchet' c connId force = withConnLock c connId "synchronizeRatchet" ackQueueMessage :: AgentMonad m => AgentClient -> RcvQueue -> SMP.MsgId -> m () ackQueueMessage c rq srvMsgId = - sendAck c rq srvMsgId `catchError` \case + sendAck c rq srvMsgId `catchAgentError` \case SMP SMP.NO_MSG -> pure () e -> throwError e @@ -1511,7 +1510,7 @@ registerNtfToken' c suppliedDeviceToken suppliedNtfMode = replaceToken :: NtfTokenId -> m NtfTknStatus replaceToken tknId = do ns <- asks ntfSupervisor - tryReplace ns `catchError` \e -> + tryReplace ns `catchAgentError` \e -> if temporaryOrHostError e then throwError e else do @@ -1618,7 +1617,7 @@ deleteToken_ c tkn@NtfToken {ntfTokenId, ntfTknStatus} = do let ntfTknAction = Just NTADelete withStore' c $ \db -> updateNtfToken db tkn ntfTknStatus ntfTknAction atomically $ nsUpdateToken ns tkn {ntfTknStatus, ntfTknAction} - agentNtfDeleteToken c tknId tkn `catchError` \case + agentNtfDeleteToken c tknId tkn `catchAgentError` \case NTF AUTH -> pure () e -> throwError e withStore' c $ \db -> removeNtfToken db tkn @@ -1728,16 +1727,16 @@ cleanupManager c@AgentClient {subQ} = do int <- asks (cleanupInterval . config) forever $ do void . runExceptT $ do - deleteConns `catchError` (notify "" . ERR) - deleteRcvMsgHashes `catchError` (notify "" . ERR) - deleteProcessedRatchetKeyHashes `catchError` (notify "" . ERR) - deleteRcvFilesExpired `catchError` (notify "" . RFERR) - deleteRcvFilesDeleted `catchError` (notify "" . RFERR) - deleteRcvFilesTmpPaths `catchError` (notify "" . RFERR) - deleteSndFilesExpired `catchError` (notify "" . SFERR) - deleteSndFilesDeleted `catchError` (notify "" . SFERR) - deleteSndFilesPrefixPaths `catchError` (notify "" . SFERR) - deleteExpiredReplicasForDeletion `catchError` (notify "" . SFERR) + deleteConns `catchAgentError` (notify "" . ERR) + deleteRcvMsgHashes `catchAgentError` (notify "" . ERR) + deleteProcessedRatchetKeyHashes `catchAgentError` (notify "" . ERR) + deleteRcvFilesExpired `catchAgentError` (notify "" . RFERR) + deleteRcvFilesDeleted `catchAgentError` (notify "" . RFERR) + deleteRcvFilesTmpPaths `catchAgentError` (notify "" . RFERR) + deleteSndFilesExpired `catchAgentError` (notify "" . SFERR) + deleteSndFilesDeleted `catchAgentError` (notify "" . SFERR) + deleteSndFilesPrefixPaths `catchAgentError` (notify "" . SFERR) + deleteExpiredReplicasForDeletion `catchAgentError` (notify "" . SFERR) liftIO $ threadDelay' int where deleteConns = @@ -1753,33 +1752,33 @@ cleanupManager c@AgentClient {subQ} = do deleteRcvFilesExpired = do rcvFilesTTL <- asks $ rcvFilesTTL . config rcvExpired <- withStore' c (`getRcvFilesExpired` rcvFilesTTL) - forM_ rcvExpired $ \(dbId, entId, p) -> flip catchError (notify entId . RFERR) $ do + forM_ rcvExpired $ \(dbId, entId, p) -> flip catchAgentError (notify entId . RFERR) $ do removePath =<< toFSFilePath p withStore' c (`deleteRcvFile'` dbId) deleteRcvFilesDeleted = do rcvDeleted <- withStore' c getCleanupRcvFilesDeleted - forM_ rcvDeleted $ \(dbId, entId, p) -> flip catchError (notify entId . RFERR) $ do + forM_ rcvDeleted $ \(dbId, entId, p) -> flip catchAgentError (notify entId . RFERR) $ do removePath =<< toFSFilePath p withStore' c (`deleteRcvFile'` dbId) deleteRcvFilesTmpPaths = do rcvTmpPaths <- withStore' c getCleanupRcvFilesTmpPaths - forM_ rcvTmpPaths $ \(dbId, entId, p) -> flip catchError (notify entId . RFERR) $ do + forM_ rcvTmpPaths $ \(dbId, entId, p) -> flip catchAgentError (notify entId . RFERR) $ do removePath =<< toFSFilePath p withStore' c (`updateRcvFileNoTmpPath` dbId) deleteSndFilesExpired = do sndFilesTTL <- asks $ sndFilesTTL . config sndExpired <- withStore' c (`getSndFilesExpired` sndFilesTTL) - forM_ sndExpired $ \(dbId, entId, p) -> flip catchError (notify entId . SFERR) $ do + forM_ sndExpired $ \(dbId, entId, p) -> flip catchAgentError (notify entId . SFERR) $ do forM_ p $ removePath <=< toFSFilePath withStore' c (`deleteSndFile'` dbId) deleteSndFilesDeleted = do sndDeleted <- withStore' c getCleanupSndFilesDeleted - forM_ sndDeleted $ \(dbId, entId, p) -> flip catchError (notify entId . SFERR) $ do + forM_ sndDeleted $ \(dbId, entId, p) -> flip catchAgentError (notify entId . SFERR) $ do forM_ p $ removePath <=< toFSFilePath withStore' c (`deleteSndFile'` dbId) deleteSndFilesPrefixPaths = do sndPrefixPaths <- withStore' c getCleanupSndFilesPrefixPaths - forM_ sndPrefixPaths $ \(dbId, entId, p) -> flip catchError (notify entId . SFERR) $ do + forM_ sndPrefixPaths $ \(dbId, entId, p) -> flip catchAgentError (notify entId . SFERR) $ do removePath =<< toFSFilePath p withStore' c (`updateSndFileNoPrefixPath` dbId) deleteExpiredReplicasForDeletion = do @@ -1944,7 +1943,7 @@ processSMPTransmission c@AgentClient {smpClients, subQ} (tSess@(_, srv, _), v, s ackDel :: InternalId -> m () ackDel = enqueueCmd . ICAckDel rId srvMsgId handleNotifyAck :: m () -> m () - handleNotifyAck m = m `catchError` \e -> notify (ERR e) >> ack + handleNotifyAck m = m `catchAgentError` \e -> notify (ERR e) >> ack SMP.END -> atomically (TM.lookup tSess smpClients $>>= tryReadTMVar >>= processEND) >>= logServer "<--" c srv rId @@ -2066,7 +2065,7 @@ processSMPTransmission c@AgentClient {smpClients, subQ} (tSess@(_, srv, _), v, s RcvConnection {} -> do AcceptedConfirmation {ownConnInfo} <- withStore c (`getAcceptedConfirmation` connId) let cData' = toConnData conn' - connectReplyQueues c cData' ownConnInfo smpQueues `catchError` (notify . ERR) + connectReplyQueues c cData' ownConnInfo smpQueues `catchAgentError` (notify . ERR) _ -> prohibited continueSending :: (SMPServer, SMP.SenderId) -> Connection 'CDuplex -> m () diff --git a/src/Simplex/Messaging/Agent/Client.hs b/src/Simplex/Messaging/Agent/Client.hs index d4289749b..1d33fa1bc 100644 --- a/src/Simplex/Messaging/Agent/Client.hs +++ b/src/Simplex/Messaging/Agent/Client.hs @@ -443,7 +443,7 @@ reconnectServer c tSess = newAsyncAction tryReconnectSMPClient $ reconnections c tryReconnectSMPClient aId = do ri <- asks $ reconnectInterval . config withRetryInterval ri $ \_ loop -> - reconnectSMPClient c tSess `catchError` const loop + reconnectSMPClient c tSess `catchAgentError` const loop atomically . removeAsyncAction aId $ reconnections c reconnectSMPClient :: forall m. AgentMonad m => AgentClient -> SMPTransportSession -> m () @@ -640,7 +640,7 @@ withLockMap_ locks key = withGetLock $ TM.lookup key locks >>= maybe newLock pur withClient_ :: forall a m err msg. (AgentMonad m, ProtocolServerClient err msg) => AgentClient -> TransportSession msg -> ByteString -> (Client msg -> m a) -> m a withClient_ c tSess@(userId, srv, _) statCmd action = do cl <- getProtocolServerClient c tSess - (action cl <* stat cl "OK") `catchError` logServerError cl + (action cl <* stat cl "OK") `catchAgentError` logServerError cl where stat cl = liftIO . incClientStat c userId cl statCmd logServerError :: Client msg -> AgentErrorType -> m a diff --git a/src/Simplex/Messaging/Agent/Env/SQLite.hs b/src/Simplex/Messaging/Agent/Env/SQLite.hs index e199f0b79..ad1d882bb 100644 --- a/src/Simplex/Messaging/Agent/Env/SQLite.hs +++ b/src/Simplex/Messaging/Agent/Env/SQLite.hs @@ -6,6 +6,7 @@ {-# LANGUAGE NamedFieldPuns #-} {-# LANGUAGE NumericUnderscores #-} {-# LANGUAGE RankNTypes #-} +{-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TypeApplications #-} {-# OPTIONS_GHC -fno-warn-unticked-promoted-constructors #-} @@ -17,6 +18,9 @@ module Simplex.Messaging.Agent.Env.SQLite NetworkConfig (..), defaultAgentConfig, defaultReconnectInterval, + tryAgentError, + catchAgentError, + agentFinally, Env (..), newSMPAgentEnv, createAgentStore, @@ -52,9 +56,10 @@ import Simplex.Messaging.TMap (TMap) import qualified Simplex.Messaging.TMap as TM import Simplex.Messaging.Transport (TLS, Transport (..)) import Simplex.Messaging.Transport.Client (defaultSMPPort) +import Simplex.Messaging.Util (allFinally, catchAllErrors, tryAllErrors) import Simplex.Messaging.Version import System.Random (StdGen, newStdGen) -import UnliftIO (Async) +import UnliftIO (Async, SomeException) import UnliftIO.STM type AgentMonad' m = (MonadUnliftIO m, MonadReader Env m) @@ -225,3 +230,19 @@ newXFTPAgent = do xftpSndWorkers <- TM.empty xftpDelWorkers <- TM.empty pure XFTPAgent {xftpWorkDir, xftpRcvWorkers, xftpSndWorkers, xftpDelWorkers} + +tryAgentError :: AgentMonad m => m a -> m (Either AgentErrorType a) +tryAgentError = tryAllErrors mkInternal +{-# INLINE tryAgentError #-} + +catchAgentError :: AgentMonad m => m a -> (AgentErrorType -> m a) -> m a +catchAgentError = catchAllErrors mkInternal +{-# INLINE catchAgentError #-} + +agentFinally :: AgentMonad m => m a -> m a -> m a +agentFinally = allFinally mkInternal +{-# INLINE agentFinally #-} + +mkInternal :: SomeException -> AgentErrorType +mkInternal = INTERNAL . show +{-# INLINE mkInternal #-} diff --git a/src/Simplex/Messaging/Agent/NtfSubSupervisor.hs b/src/Simplex/Messaging/Agent/NtfSubSupervisor.hs index 186b2cbb0..8e4603683 100644 --- a/src/Simplex/Messaging/Agent/NtfSubSupervisor.hs +++ b/src/Simplex/Messaging/Agent/NtfSubSupervisor.hs @@ -147,7 +147,7 @@ processNtfSub c (connId, cmd) = do atomically (TM.lookup srv ws) >>= \case Nothing -> do doWork <- newTMVarIO () - worker <- async $ runWorker c srv doWork `E.finally` atomically (TM.delete srv ws) + worker <- async $ runWorker c srv doWork `agentFinally` atomically (TM.delete srv ws) atomically $ TM.insert srv (doWork, worker) ws Just (doWork, _) -> void . atomically $ tryPutTMVar doWork () @@ -173,7 +173,7 @@ runNtfWorker c srv doWork = do ri <- asks $ reconnectInterval . config withRetryInterval ri $ \_ loop -> processAction a - `catchError` retryOnError c "NtfWorker" loop (workerInternalError c connId . show) + `catchAgentError` retryOnError c "NtfWorker" loop (workerInternalError c connId . show) noWorkToDo = void . atomically $ tryTakeTMVar doWork processAction :: (NtfSubscription, NtfSubNTFAction, NtfActionTs) -> m () processAction (sub@NtfSubscription {connId, smpServer, ntfSubId}, action, actionTs) = do @@ -213,7 +213,7 @@ runNtfWorker c srv doWork = do NSADelete -> case ntfSubId of Just nSubId -> (getNtfToken >>= mapM_ (agentNtfDeleteSubscription c nSubId)) - `E.finally` continueDeletion + `agentFinally` continueDeletion _ -> continueDeletion where continueDeletion = do @@ -224,7 +224,7 @@ runNtfWorker c srv doWork = do NSARotate -> case ntfSubId of Just nSubId -> (getNtfToken >>= mapM_ (agentNtfDeleteSubscription c nSubId)) - `E.finally` deleteCreate + `agentFinally` deleteCreate _ -> deleteCreate where deleteCreate = do @@ -257,7 +257,7 @@ runNtfSMPWorker c srv doWork = do ri <- asks $ reconnectInterval . config withRetryInterval ri $ \_ loop -> processAction a - `catchError` retryOnError c "NtfSMPWorker" loop (workerInternalError c connId . show) + `catchAgentError` retryOnError c "NtfSMPWorker" loop (workerInternalError c connId . show) noWorkToDo = void . atomically $ tryTakeTMVar doWork processAction :: (NtfSubscription, NtfSubSMPAction, NtfActionTs) -> m () processAction (sub@NtfSubscription {connId, ntfServer}, smpAction, actionTs) = do diff --git a/src/Simplex/Messaging/Util.hs b/src/Simplex/Messaging/Util.hs index ec3678b39..aecb59aa3 100644 --- a/src/Simplex/Messaging/Util.hs +++ b/src/Simplex/Messaging/Util.hs @@ -1,4 +1,3 @@ -{-# LANGUAGE NumericUnderscores #-} {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE ScopedTypeVariables #-} @@ -13,12 +12,13 @@ import Data.Bifunctor (first) import Data.ByteString.Char8 (ByteString) import qualified Data.ByteString.Char8 as B import Data.Int (Int64) +import Data.List (groupBy, sortOn) import Data.Text (Text) import qualified Data.Text as T import Data.Text.Encoding (decodeUtf8With) import Data.Time (NominalDiffTime) import UnliftIO.Async -import Data.List (groupBy, sortOn) +import qualified UnliftIO.Exception as UE raceAny_ :: MonadUnliftIO m => [m a] -> m () raceAny_ = r [] @@ -99,17 +99,34 @@ catchAll_ :: IO a -> IO a -> IO a catchAll_ a = catchAll a . const {-# INLINE catchAll_ #-} +tryAllErrors :: (MonadUnliftIO m, MonadError e m) => (E.SomeException -> e) -> m a -> m (Either e a) +tryAllErrors err action = tryError action `UE.catch` (pure . Left . err) +{-# INLINE tryAllErrors #-} + +catchAllErrors :: (MonadUnliftIO m, MonadError e m) => (E.SomeException -> e) -> m a -> (e -> m a) -> m a +catchAllErrors err action handle = tryAllErrors err action >>= either handle pure +{-# INLINE catchAllErrors #-} + +catchThrow :: (MonadUnliftIO m, MonadError e m) => m a -> (E.SomeException -> e) -> m a +catchThrow action err = catchAllErrors err action throwError +{-# INLINE catchThrow #-} + +allFinally :: (MonadUnliftIO m, MonadError e m) => (E.SomeException -> e) -> m a -> m a -> m a +allFinally err action final = tryAllErrors err action >>= either (\e -> final >> throwError e) (const final) +{-# INLINE allFinally #-} + eitherToMaybe :: Either a b -> Maybe b eitherToMaybe = either (const Nothing) Just {-# INLINE eitherToMaybe #-} groupOn :: Eq k => (a -> k) -> [a] -> [[a]] groupOn = groupBy . eqOn - -- it is equivalent to groupBy ((==) `on` f), - -- but it redefines `on` to avoid duplicate computation for most values. - -- source: https://hackage.haskell.org/package/extra-1.7.13/docs/src/Data.List.Extra.html#groupOn - -- the on2 in this package is specialized to only use `==` as the function, `eqOn f` is equivalent to `(==) `on` f` - where eqOn f = \x -> let fx = f x in \y -> fx == f y + -- it is equivalent to groupBy ((==) `on` f), + -- but it redefines `on` to avoid duplicate computation for most values. + -- source: https://hackage.haskell.org/package/extra-1.7.13/docs/src/Data.List.Extra.html#groupOn + -- the on2 in this package is specialized to only use `==` as the function, `eqOn f` is equivalent to `(==) `on` f` + where + eqOn f = \x -> let fx = f x in \y -> fx == f y groupAllOn :: Ord k => (a -> k) -> [a] -> [[a]] groupAllOn f = groupOn f . sortOn f diff --git a/tests/CoreTests/UtilTests.hs b/tests/CoreTests/UtilTests.hs new file mode 100644 index 000000000..3da316ecc --- /dev/null +++ b/tests/CoreTests/UtilTests.hs @@ -0,0 +1,110 @@ +{-# LANGUAGE ScopedTypeVariables #-} + +module CoreTests.UtilTests where + +import Control.Exception (Exception, SomeException, throwIO) +import Control.Monad.Except +import Data.IORef +import Simplex.Messaging.Util +import Simplex.Messaging.Client.Agent () +import Test.Hspec +import qualified UnliftIO.Exception as UE + +utilTests :: Spec +utilTests = do + describe "lifted try, catch and finally problems" $ do + describe "try" $ do + it "lifted try does not catch errors" $ do + runExceptT (UE.try throwTestError >>= either handleCatch pure) `shouldReturn` Left (TestError "error") + runExceptT (UE.try throwTestException >>= either handleCatch pure) `shouldThrow` (\(e :: IOError) -> show e == "user error (error)") + it "lifted try with SomeException catches all errors but wraps ExceptT errors" $ do + runExceptT (UE.try throwTestError >>= either handleException pure) `shouldReturn` Right "caught InternalException {unInternalException = TestError \"error\"}" + runExceptT (UE.try throwTestException >>= either handleException pure) `shouldReturn` Right "caught user error (error)" + describe "catch" $ do + it "lifted catch does not catch" $ do + runExceptT (throwTestError `UE.catch` handleCatch) `shouldReturn` Left (TestError "error") + runExceptT (throwTestException `UE.catch` handleCatch) `shouldThrow` (\(e :: IOError) -> show e == "user error (error)") + it "lifted catch of SomeException catches all errors but wraps ExceptT errors" $ do + runExceptT (throwTestError `UE.catch` handleException) `shouldReturn` Right "caught InternalException {unInternalException = TestError \"error\"}" + runExceptT (throwTestException `UE.catch` handleException) `shouldReturn` Right "caught user error (error)" + describe "finally" $ do + it "lifted finally executes final action and stays in ExceptT monad" $ withFinal $ \final -> + runExceptT (throwTestError `UE.finally` final) `shouldReturn` Left (TestError "error") + it "lifted finally executes final action but throws exception" $ withFinal $ \final -> + runExceptT (throwTestException `UE.finally` final) `shouldThrow` (\(e :: IOError) -> show e == "user error (error)") + describe "tryAllErrors" $ do + it "should return ExceptT error as Left" $ + runExceptT (tryAllErrors testErr throwTestError) `shouldReturn` Right (Left (TestError "error")) + it "should return SomeException as Left" $ + runExceptT (tryAllErrors testErr throwTestException) `shouldReturn` Right (Left (TestException "user error (error)")) + it "should return no errors as Right" $ + runExceptT (tryAllErrors testErr noErrors) `shouldReturn` Right (Right "no errors") + describe "tryAllErrors specialized as tryTestError" $ do + let tryTestError = tryAllErrors testErr + it "should return ExceptT error as Left" $ + runExceptT (tryTestError throwTestError) `shouldReturn` Right (Left (TestError "error")) + it "should return SomeException as Left" $ + runExceptT (tryTestError throwTestException) `shouldReturn` Right (Left (TestException "user error (error)")) + it "should return no errors as Right" $ + runExceptT (tryTestError noErrors) `shouldReturn` Right (Right "no errors") + describe "catchAllErrors" $ do + it "should catch ExceptT error" $ + runExceptT (catchAllErrors testErr throwTestError handleCatch) `shouldReturn` Right "caught TestError \"error\"" + it "should catch SomeException" $ + runExceptT (catchAllErrors testErr throwTestException handleCatch) `shouldReturn` Right "caught TestException \"user error (error)\"" + it "should not throw if there are no errors" $ + runExceptT (catchAllErrors testErr noErrors throwError) `shouldReturn` Right "no errors" + describe "catchAllErrors specialized as catchTestError" $ do + let catchTestError = catchAllErrors testErr + it "should catch ExceptT error" $ + runExceptT (throwTestError `catchTestError` handleCatch) `shouldReturn` Right "caught TestError \"error\"" + it "should catch SomeException" $ + runExceptT (throwTestException `catchTestError` handleCatch) `shouldReturn` Right "caught TestException \"user error (error)\"" + it "should not throw if there are no errors" $ + runExceptT (noErrors `catchTestError` throwError) `shouldReturn` Right "no errors" + describe "catchThrow" $ do + it "should re-throw ExceptT error" $ + runExceptT (throwTestError `catchThrow` testErr) `shouldReturn` Left (TestError "error") + it "should catch SomeException and throw as ExceptT error" $ + runExceptT (throwTestException `catchThrow` testErr) `shouldReturn` Left (TestException "user error (error)") + it "should not throw if there are no exceptions" $ + runExceptT (noErrors `catchThrow` testErr) `shouldReturn` Right "no errors" + describe "allFinally should run final action" $ do + it "then throw ExceptT error" $ withFinal $ \final -> + runExceptT (allFinally testErr throwTestError final) `shouldReturn` Left (TestError "error") + it "then throw SomeException as ExceptT error" $ withFinal $ \final -> + runExceptT (allFinally testErr throwTestException final) `shouldReturn` Left (TestException "user error (error)") + it "and should not throw if there are no exceptions" $ withFinal $ \final -> + runExceptT (allFinally testErr noErrors final) `shouldReturn` Right "final" + describe "allFinally specialized as testFinally should run final action" $ do + let testFinally = allFinally testErr + it "then throw ExceptT error" $ withFinal $ \final -> + runExceptT (throwTestError `testFinally` final) `shouldReturn` Left (TestError "error") + it "then throw SomeException as ExceptT error" $ withFinal $ \final -> + runExceptT (throwTestException `testFinally` final) `shouldReturn` Left (TestException "user error (error)") + it "and should not throw if there are no exceptions" $ withFinal $ \final -> + runExceptT (noErrors `testFinally` final) `shouldReturn` Right "final" + where + throwTestError :: ExceptT TestError IO String + throwTestError = throwError $ TestError "error" + throwTestException :: ExceptT TestError IO String + throwTestException = liftIO $ throwIO $ userError "error" + noErrors :: ExceptT TestError IO String + noErrors = pure "no errors" + testErr :: SomeException -> TestError + testErr = TestException . show + handleCatch :: TestError -> ExceptT TestError IO String + handleCatch e = pure $ "caught " <> show e + handleException :: SomeException -> ExceptT TestError IO String + handleException e = pure $ "caught " <> show e + withFinal :: (ExceptT TestError IO String -> IO ()) -> IO () + withFinal test = do + r <- newIORef False + let final = liftIO $ writeIORef r True >> pure "final" + test final + readIORef r `shouldReturn` True + +data TestError = TestError String | TestException String + deriving (Eq, Show) + +instance Exception TestError diff --git a/tests/Test.hs b/tests/Test.hs index d76357be4..259c3c3cb 100644 --- a/tests/Test.hs +++ b/tests/Test.hs @@ -7,6 +7,7 @@ import CoreTests.CryptoTests import CoreTests.EncodingTests import CoreTests.ProtocolErrorTests import CoreTests.RetryIntervalTests +import CoreTests.UtilTests import CoreTests.VersionRangeTests import FileDescriptionTests (fileDescriptionTests) import NtfServerTests (ntfServerTests) @@ -39,6 +40,7 @@ main = do describe "Version range" versionRangeTests describe "Encryption tests" cryptoTests describe "Retry interval tests" retryIntervalTests + describe "Util tests" utilTests describe "SMP server via TLS" $ serverTests (transport @TLS) describe "SMP server via WebSockets" $ serverTests (transport @WS) describe "Notifications server" $ ntfServerTests (transport @TLS)