From 55d31d1fefb7a87b2db4cb14201a35aae7c1648c Mon Sep 17 00:00:00 2001 From: Evgeny Poberezkin Date: Fri, 3 May 2024 00:00:17 +0100 Subject: [PATCH] move check and state update to one transaction --- src/Simplex/Messaging/Agent.hs | 3 +- src/Simplex/Messaging/Agent/Client.hs | 59 ++++++++++++++------------- 2 files changed, 31 insertions(+), 31 deletions(-) diff --git a/src/Simplex/Messaging/Agent.hs b/src/Simplex/Messaging/Agent.hs index 72d4ad3b4..011d4128c 100644 --- a/src/Simplex/Messaging/Agent.hs +++ b/src/Simplex/Messaging/Agent.hs @@ -121,7 +121,6 @@ import Control.Logger.Simple (logError, logInfo, showText) import Control.Monad import Control.Monad.Except import Control.Monad.Reader -import Control.Monad.Trans.Except (throwE) import Crypto.Random (ChaChaDRG) import qualified Data.Aeson as J import Data.Bifunctor (bimap, first, second) @@ -891,7 +890,7 @@ subscribeConnections' c connIds = do (subRs, rcvQs) = M.mapEither rcvQueueOrResult cs mapM_ (mapM_ (\(cData, sqs) -> mapM_ (lift . resumeMsgDelivery c cData) sqs) . sndQueue) cs mapM_ (resumeConnCmds c) $ M.keys cs - rcvRs <- lift $ connResults . snd <$> subscribeQueues c (concat $ M.elems rcvQs) + rcvRs <- lift $ connResults . fst <$> subscribeQueues c (concat $ M.elems rcvQs) ns <- asks ntfSupervisor tkn <- readTVarIO (ntfTkn ns) when (instantNotifications tkn) . void . lift . forkIO . void . runExceptT $ sendNtfCreate ns rcvRs conns diff --git a/src/Simplex/Messaging/Agent/Client.hs b/src/Simplex/Messaging/Agent/Client.hs index a9c6d6745..73c67dddb 100644 --- a/src/Simplex/Messaging/Agent/Client.hs +++ b/src/Simplex/Messaging/Agent/Client.hs @@ -155,7 +155,7 @@ import Data.Bifunctor (bimap, first, second) import Data.ByteString.Base64 import Data.ByteString.Char8 (ByteString) import qualified Data.ByteString.Char8 as B -import Data.Either (lefts, partitionEithers) +import Data.Either (partitionEithers) import Data.Functor (($>)) import Data.Int (Int64) import Data.List (deleteFirstsBy, foldl', partition, (\\)) @@ -643,7 +643,7 @@ reconnectSMPClient tc c tSess@(_, srv, _) qs = do resubscribe :: AM () resubscribe = do cs <- readTVarIO $ RQ.getConnections $ activeSubs c - (sessions, rs) <- lift . subscribeQueues c $ L.toList qs + (rs, sessId_) <- lift . subscribeQueues c $ L.toList qs let (errs, okConns) = partitionEithers $ map (\(RcvQueue {connId}, r) -> bimap (connId,) (const connId) r) rs liftIO $ do let conns = filter (`M.notMember` cs) okConns @@ -652,8 +652,8 @@ reconnectSMPClient tc c tSess@(_, srv, _) qs = do liftIO $ mapM_ (\(connId, e) -> notifySub connId $ ERR e) finalErrs forM_ (listToMaybe tempErrs) $ \(_, err) -> do when (null okConns && M.null cs && null finalErrs) . liftIO $ - forM_ (listToMaybe sessions) $ \sessId -> do - -- only close the client session that was used to subscribe + forM_ sessId_ $ \sessId -> do + -- We only close the client session that was used to subscribe. v_ <- atomically $ ifM (activeClientSession c tSess sessId) (TM.lookupDelete tSess $ smpClients c) (pure Nothing) mapM_ (closeClient_ c) v_ throwError err @@ -1136,15 +1136,13 @@ newRcvQueue c userId connId (ProtoServerWithAuth srv auth) vRange subMode = do qUri = SMPQueueUri vRange $ SMPQueueAddress srv sndId e2eDhKey pure (rq, qUri, tSess, sessId) -processSubResult :: AgentClient -> RcvQueue -> Either SMPClientError () -> IO (Either SMPClientError ()) -processSubResult c rq r = do - case r of - Left e -> - unless (temporaryClientError e) . atomically $ do - RQ.deleteQueue rq (pendingSubs c) - TM.insert (RQ.qKey rq) e (removedSubs c) - _ -> atomically $ addSubscription c rq - pure r +processSubResult :: AgentClient -> RcvQueue -> Either SMPClientError () -> STM () +processSubResult c rq = \case + Left e -> + unless (temporaryClientError e) $ do + RQ.deleteQueue rq (pendingSubs c) + TM.insert (RQ.qKey rq) e (removedSubs c) + Right () -> addSubscription c rq temporaryAgentError :: AgentErrorType -> Bool temporaryAgentError = \case @@ -1161,7 +1159,7 @@ temporaryOrHostError = \case {-# INLINE temporaryOrHostError #-} -- | Subscribe to queues. The list of results can have a different order. -subscribeQueues :: AgentClient -> [RcvQueue] -> AM' ([SessionId], [(RcvQueue, Either AgentErrorType ())]) +subscribeQueues :: AgentClient -> [RcvQueue] -> AM' ([(RcvQueue, Either AgentErrorType ())], Maybe SessionId) subscribeQueues c qs = do (errs, qs') <- partitionEithers <$> mapM checkQueue qs atomically $ do @@ -1169,27 +1167,30 @@ subscribeQueues c qs = do RQ.batchAddQueues (pendingSubs c) qs' env <- ask -- only "checked" queues are subscribed - sessionsUsed <- newTVarIO [] - rs <- sendTSessionBatches "SUB" 90 id (subscribeQueues_ env sessionsUsed) c qs' - sessionsUsed' <- readTVarIO sessionsUsed - pure (sessionsUsed', errs <> rs) + session <- newTVarIO Nothing + rs <- sendTSessionBatches "SUB" 90 id (subscribeQueues_ env session) c qs' + (errs <> rs,) <$> readTVarIO session where checkQueue rq = do prohibited <- atomically $ hasGetLock c rq pure $ if prohibited then Left (rq, Left $ CMD PROHIBITED) else Right rq - subscribeQueues_ :: Env -> TVar [SessionId] -> SMPClient -> NonEmpty RcvQueue -> IO (BatchResponses SMPClientError ()) - subscribeQueues_ env sessionsUsed smp qs' = do + subscribeQueues_ :: Env -> TVar (Maybe SessionId) -> SMPClient -> NonEmpty RcvQueue -> IO (BatchResponses SMPClientError ()) + subscribeQueues_ env session smp qs' = do rs <- sendBatch subscribeSMPQueues smp qs' - let sessId = sessionId $ thParams smp - active <- atomically $ activeClientSession c (transportSession' smp) sessId - if active - then do - atomically $ modifyTVar' sessionsUsed (sessId :) - mapM_ (uncurry $ processSubResult c) rs - when (any temporaryClientError . lefts . map snd $ L.toList rs) $ - runReaderT (resubscribeSMPSession c $ transportSession' smp) env - else logWarn "subcription batch result for replaced SMP client, processing skipped" + let tSess = transportSession' smp + sessId = sessionId $ thParams smp + active <- + atomically $ + ifM + (activeClientSession c tSess sessId) + (writeTVar session (Just sessId) >> mapM_ (uncurry $ processSubResult c) rs $> True) + (pure False) + when (active || hasTempErrors rs) $ + resubscribeSMPSession c tSess `runReaderT` env + unless active $ logWarn "subcription batch result for replaced SMP client, resubscribing" pure rs + where + hasTempErrors = any (either temporaryClientError (const False) . snd) activeClientSession :: AgentClient -> SMPTransportSession -> SessionId -> STM Bool activeClientSession c tSess sessId = sameSess <$> tryReadSessVar tSess (smpClients c)