From a0b35cec4f39e248eae45799667e0b483b94e831 Mon Sep 17 00:00:00 2001 From: Evgeny Poberezkin <2769109+epoberezkin@users.noreply.github.com> Date: Tue, 16 Jan 2024 13:45:51 +0000 Subject: [PATCH] agent: fix potential race when good client can be removed instead of bad for the same transport session (#967) * agent: fix potential race when good client can be removed instead of bad for the same transport session * tryAgentError * case --- src/Simplex/FileTransfer/Client.hs | 3 - src/Simplex/Messaging/Agent.hs | 2 +- src/Simplex/Messaging/Agent/Client.hs | 134 ++++++++++++-------------- 3 files changed, 62 insertions(+), 77 deletions(-) diff --git a/src/Simplex/FileTransfer/Client.hs b/src/Simplex/FileTransfer/Client.hs index 9109de789..9489f52c1 100644 --- a/src/Simplex/FileTransfer/Client.hs +++ b/src/Simplex/FileTransfer/Client.hs @@ -116,9 +116,6 @@ xftpTransportHost XFTPClient {http2Client = HTTP2Client {client_ = HClient {host xftpSessionTs :: XFTPClient -> UTCTime xftpSessionTs = sessionTs . http2Client -xftpSessionId :: XFTPClient -> ByteString -xftpSessionId = sessionId . http2Client - xftpHTTP2Config :: TransportClientConfig -> XFTPClientConfig -> HTTP2ClientConfig xftpHTTP2Config transportConfig XFTPClientConfig {xftpNetworkConfig = NetworkConfig {tcpConnectTimeout}} = defaultHTTP2ClientConfig diff --git a/src/Simplex/Messaging/Agent.hs b/src/Simplex/Messaging/Agent.hs index f552c59d8..bcec478fd 100644 --- a/src/Simplex/Messaging/Agent.hs +++ b/src/Simplex/Messaging/Agent.hs @@ -2047,7 +2047,7 @@ processSMPTransmission c@AgentClient {smpClients, subQ} (tSess@(_, srv, _), v, s handleNotifyAck :: m () -> m () handleNotifyAck m = m `catchAgentError` \e -> notify (ERR e) >> ack SMP.END -> - atomically (TM.lookup tSess smpClients $>>= tryReadTMVar >>= processEND) + atomically (TM.lookup tSess smpClients $>>= (tryReadTMVar . sessionVar) >>= processEND) >>= logServer "<--" c srv rId where processEND = \case diff --git a/src/Simplex/Messaging/Agent/Client.hs b/src/Simplex/Messaging/Agent/Client.hs index 8f04f8237..305f39c26 100644 --- a/src/Simplex/Messaging/Agent/Client.hs +++ b/src/Simplex/Messaging/Agent/Client.hs @@ -78,6 +78,7 @@ module Simplex.Messaging.Agent.Client agentDRG, getAgentSubscriptions, Worker (..), + SessionVar (..), SubscriptionsInfo (..), SubInfo (..), AgentOperation (..), @@ -222,13 +223,18 @@ import UnliftIO.Exception (bracket) import qualified UnliftIO.Exception as E import UnliftIO.STM -type ClientVar msg = TMVar (Either AgentErrorType (Client msg)) +data SessionVar a = SessionVar + { sessionVar :: TMVar a, + sessionVarId :: Int + } + +type ClientVar msg = SessionVar (Either AgentErrorType (Client msg)) type SMPClientVar = ClientVar SMP.BrokerMsg type NtfClientVar = ClientVar NtfResponse -type XFTPClientVar = TMVar (Either AgentErrorType XFTPClient) +type XFTPClientVar = ClientVar FileResponse type SMPTransportSession = TransportSession SMP.BrokerMsg @@ -270,18 +276,13 @@ data AgentClient = AgentClient -- lock to prevent concurrency between periodic and async connection deletions deleteLock :: Lock, -- smpSubWorkers for SMP servers sessions - smpSubWorkers :: TMap SMPTransportSession (TMVar SubWorker), + smpSubWorkers :: TMap SMPTransportSession (SessionVar (Async ())), asyncClients :: TAsyncs, agentStats :: TMap AgentStatsKey (TVar Int), clientId :: Int, agentEnv :: Env } -data SubWorker = SubWorker - { subWorkerId :: Int, - subWorkerAsync :: Async () - } - getAgentWorker :: (AgentMonad' m, Ord k, Show k) => String -> Bool -> AgentClient -> k -> TMap k Worker -> (Worker -> ExceptT AgentErrorType m ()) -> m Worker getAgentWorker = getAgentWorker' id pure @@ -470,7 +471,6 @@ class (Encoding err, Show err) => ProtocolServerClient err msg | msg -> err wher clientServer :: Client msg -> String clientTransportHost :: Client msg -> TransportHost clientSessionTs :: Client msg -> UTCTime - clientSessionId :: Client msg -> ByteString instance ProtocolServerClient ErrorType BrokerMsg where type Client BrokerMsg = ProtocolClient ErrorType BrokerMsg @@ -480,7 +480,6 @@ instance ProtocolServerClient ErrorType BrokerMsg where clientServer = protocolClientServer clientTransportHost = transportHost' clientSessionTs = sessionTs - clientSessionId = sessionId instance ProtocolServerClient ErrorType NtfResponse where type Client NtfResponse = ProtocolClient ErrorType NtfResponse @@ -490,7 +489,6 @@ instance ProtocolServerClient ErrorType NtfResponse where clientServer = protocolClientServer clientTransportHost = transportHost' clientSessionTs = sessionTs - clientSessionId = sessionId instance ProtocolServerClient XFTPErrorType FileResponse where type Client FileResponse = XFTPClient @@ -500,29 +498,28 @@ instance ProtocolServerClient XFTPErrorType FileResponse where clientServer = X.xftpClientServer clientTransportHost = X.xftpTransportHost clientSessionTs = X.xftpSessionTs - clientSessionId = X.xftpSessionId getSMPServerClient :: forall m. AgentMonad m => AgentClient -> SMPTransportSession -> m SMPClient getSMPServerClient c@AgentClient {active, smpClients, msgQ} tSess@(userId, srv, _) = do unlessM (readTVarIO active) . throwError $ INACTIVE - atomically (getTSessVar tSess smpClients) + atomically (getTSessVar c tSess smpClients) >>= either newClient (waitForProtocolClient c tSess) where newClient = newProtocolClient c tSess smpClients connectClient resubscribeSMPSession - connectClient :: m SMPClient - connectClient = do + connectClient :: SMPClientVar -> m SMPClient + connectClient v = do cfg <- getClientConfig c smpCfg u <- askUnliftIO - liftEitherError (protocolClientError SMP $ B.unpack $ strEncode srv) (getProtocolClient tSess cfg (Just msgQ) $ clientDisconnected u) + liftEitherError (protocolClientError SMP $ B.unpack $ strEncode srv) (getProtocolClient tSess cfg (Just msgQ) $ clientDisconnected u v) - clientDisconnected :: UnliftIO m -> SMPClient -> IO () - clientDisconnected u client = do + clientDisconnected :: UnliftIO m -> SMPClientVar -> SMPClient -> IO () + clientDisconnected u v client = do removeClientAndSubs >>= serverDown logInfo . decodeUtf8 $ "Agent disconnected from " <> showServer srv where removeClientAndSubs :: IO ([RcvQueue], [ConnId]) removeClientAndSubs = atomically $ do - removeClientVar client tSess smpClients + removeTSessVar v tSess smpClients qs <- RQ.getDelSessQueues tSess $ activeSubs c mapM_ (`RQ.addQueue` pendingSubs c) qs let cs = S.fromList $ map qConnId qs @@ -543,12 +540,11 @@ getSMPServerClient c@AgentClient {active, smpClients, msgQ} tSess@(userId, srv, resubscribeSMPSession :: AgentMonad' m => AgentClient -> SMPTransportSession -> m () resubscribeSMPSession c@AgentClient {smpSubWorkers} tSess = - atomically (getTSessVar tSess smpSubWorkers) >>= either newSubWorker (\_ -> pure ()) + atomically (getTSessVar c tSess smpSubWorkers) >>= either newSubWorker (\_ -> pure ()) where newSubWorker v = do - subWorkerId <- atomically $ stateTVar (workerSeq c) $ \next -> (next, next + 1) - subWorkerAsync <- async $ void (E.tryAny runSubWorker) >> atomically (cleanup v subWorkerId) - atomically $ putTMVar v SubWorker {subWorkerId, subWorkerAsync} + a <- async $ void (E.tryAny runSubWorker) >> atomically (cleanup v) + atomically $ putTMVar (sessionVar v) a runSubWorker = do ri <- asks $ reconnectInterval . config timeoutCounts <- newTVarIO 0 @@ -557,12 +553,12 @@ resubscribeSMPSession c@AgentClient {smpSubWorkers} tSess = forM_ (L.nonEmpty pending) $ \qs -> do void . tryAgentError' $ reconnectSMPClient timeoutCounts c tSess qs loop - cleanup :: TMVar SubWorker -> Int -> STM () - cleanup v swId = do + cleanup :: SessionVar (Async ()) -> STM () + cleanup v = do -- Here we wait until TMVar is not empty to prevent worker cleanup happening before worker is added to TMVar. -- Not waiting may result in terminated worker remaining in the map. - whenM (isEmptyTMVar v) retry - removeTSessVar ((swId ==) . subWorkerId) tSess smpSubWorkers + whenM (isEmptyTMVar $ sessionVar v) retry + removeTSessVar v tSess smpSubWorkers reconnectSMPClient :: forall m. AgentMonad m => TVar Int -> AgentClient -> SMPTransportSession -> NonEmpty RcvQueue -> m () reconnectSMPClient tc c tSess@(_, srv, _) qs = do @@ -598,19 +594,19 @@ reconnectSMPClient tc c tSess@(_, srv, _) qs = do getNtfServerClient :: forall m. AgentMonad m => AgentClient -> NtfTransportSession -> m NtfClient getNtfServerClient c@AgentClient {active, ntfClients} tSess@(userId, srv, _) = do unlessM (readTVarIO active) . throwError $ INACTIVE - atomically (getTSessVar tSess ntfClients) + atomically (getTSessVar c tSess ntfClients) >>= either (newProtocolClient c tSess ntfClients connectClient $ \_ _ -> pure ()) (waitForProtocolClient c tSess) where - connectClient :: m NtfClient - connectClient = do + connectClient :: NtfClientVar -> m NtfClient + connectClient v = do cfg <- getClientConfig c ntfCfg - liftEitherError (protocolClientError NTF $ B.unpack $ strEncode srv) (getProtocolClient tSess cfg Nothing clientDisconnected) + liftEitherError (protocolClientError NTF $ B.unpack $ strEncode srv) (getProtocolClient tSess cfg Nothing $ clientDisconnected v) - clientDisconnected :: NtfClient -> IO () - clientDisconnected client = do - atomically $ removeClientVar client tSess ntfClients + clientDisconnected :: NtfClientVar -> NtfClient -> IO () + clientDisconnected v client = do + atomically $ removeTSessVar v tSess ntfClients incClientStat c userId client "DISCONNECT" "" atomically $ writeTBQueue (subQ c) ("", "", APC SAENone $ hostEvent DISCONNECT client) logInfo . decodeUtf8 $ "Agent disconnected from " <> showServer srv @@ -618,49 +614,44 @@ getNtfServerClient c@AgentClient {active, ntfClients} tSess@(userId, srv, _) = d getXFTPServerClient :: forall m. AgentMonad m => AgentClient -> XFTPTransportSession -> m XFTPClient getXFTPServerClient c@AgentClient {active, xftpClients, useNetworkConfig} tSess@(userId, srv, _) = do unlessM (readTVarIO active) . throwError $ INACTIVE - atomically (getTSessVar tSess xftpClients) + atomically (getTSessVar c tSess xftpClients) >>= either (newProtocolClient c tSess xftpClients connectClient $ \_ _ -> pure ()) (waitForProtocolClient c tSess) where - connectClient :: m XFTPClient - connectClient = do + connectClient :: XFTPClientVar -> m XFTPClient + connectClient v = do cfg <- asks $ xftpCfg . config xftpNetworkConfig <- readTVarIO useNetworkConfig - liftEitherError (protocolClientError XFTP $ B.unpack $ strEncode srv) (X.getXFTPClient tSess cfg {xftpNetworkConfig} clientDisconnected) + liftEitherError (protocolClientError XFTP $ B.unpack $ strEncode srv) (X.getXFTPClient tSess cfg {xftpNetworkConfig} $ clientDisconnected v) - clientDisconnected :: XFTPClient -> IO () - clientDisconnected client = do - atomically $ removeClientVar client tSess xftpClients + clientDisconnected :: XFTPClientVar -> XFTPClient -> IO () + clientDisconnected v client = do + atomically $ removeTSessVar v tSess xftpClients incClientStat c userId client "DISCONNECT" "" atomically $ writeTBQueue (subQ c) ("", "", APC SAENone $ hostEvent DISCONNECT client) logInfo . decodeUtf8 $ "Agent disconnected from " <> showServer srv -getTSessVar :: forall a s. TransportSession s -> TMap (TransportSession s) (TMVar a) -> STM (Either (TMVar a) (TMVar a)) -getTSessVar tSess clients = maybe (Left <$> newClientVar) (pure . Right) =<< TM.lookup tSess clients +getTSessVar :: forall a s. AgentClient -> TransportSession s -> TMap (TransportSession s) (SessionVar a) -> STM (Either (SessionVar a) (SessionVar a)) +getTSessVar c tSess vs = maybe (Left <$> newSessionVar) (pure . Right) =<< TM.lookup tSess vs where - newClientVar :: STM (TMVar a) - newClientVar = do - var <- newEmptyTMVar - TM.insert tSess var clients - pure var + newSessionVar :: STM (SessionVar a) + newSessionVar = do + sessionVar <- newEmptyTMVar + sessionVarId <- stateTVar (workerSeq c) $ \next -> (next, next + 1) + let v = SessionVar {sessionVar, sessionVarId} + TM.insert tSess v vs + pure v -removeClientVar :: ProtocolServerClient err msg => Client msg -> TransportSession msg -> TMap (TransportSession msg) (ClientVar msg) -> STM () -removeClientVar = removeTSessVar . either (const False) . sameClient - -sameClient :: ProtocolServerClient err msg => Client msg -> Client msg -> Bool -sameClient c c' = clientSessionId c == clientSessionId c' - -removeTSessVar :: (a -> Bool) -> TransportSession msg -> TMap (TransportSession msg) (TMVar a) -> STM () -removeTSessVar same tSess vs = +removeTSessVar :: SessionVar a -> TransportSession msg -> TMap (TransportSession msg) (SessionVar a) -> STM () +removeTSessVar v tSess vs = TM.lookup tSess vs - $>>= tryReadTMVar - >>= mapM_ (\v -> when (same v) $ TM.delete tSess vs) + >>= mapM_ (\v' -> when (sessionVarId v == sessionVarId v') $ TM.delete tSess vs) waitForProtocolClient :: (AgentMonad m, ProtocolTypeI (ProtoType msg)) => AgentClient -> TransportSession msg -> ClientVar msg -> m (Client msg) -waitForProtocolClient c (_, srv, _) clientVar = do +waitForProtocolClient c (_, srv, _) v = do NetworkConfig {tcpConnectTimeout} <- readTVarIO $ useNetworkConfig c - client_ <- liftIO $ tcpConnectTimeout `timeout` atomically (readTMVar clientVar) + client_ <- liftIO $ tcpConnectTimeout `timeout` atomically (readTMVar $ sessionVar v) liftEither $ case client_ of Just (Right smpClient) -> Right smpClient Just (Left e) -> Left e @@ -673,18 +664,18 @@ newProtocolClient :: AgentClient -> TransportSession msg -> TMap (TransportSession msg) (ClientVar msg) -> - m (Client msg) -> + (ClientVar msg -> m (Client msg)) -> (AgentClient -> TransportSession msg -> m ()) -> ClientVar msg -> m (Client msg) -newProtocolClient c tSess@(userId, srv, entityId_) clients connectClient clientConnected clientVar = tryConnectClient pure tryConnectAsync +newProtocolClient c tSess@(userId, srv, entityId_) clients connectClient clientConnected v = tryConnectClient pure tryConnectAsync where tryConnectClient :: (Client msg -> m a) -> m () -> m a tryConnectClient successAction retryAction = - tryError connectClient >>= \r -> case r of + tryAgentError (connectClient v) >>= \case Right client -> do logInfo . decodeUtf8 $ "Agent connected to " <> showServer srv <> " (user " <> bshow userId <> maybe "" (" for entity " <>) entityId_ <> ")" - atomically $ putTMVar clientVar r + atomically $ putTMVar (sessionVar v) (Right client) liftIO $ incClientStat c userId client "CLIENT" "OK" atomically $ writeTBQueue (subQ c) ("", "", APC SAENone $ hostEvent CONNECT client) successAction client @@ -693,11 +684,8 @@ newProtocolClient c tSess@(userId, srv, entityId_) clients connectClient clientC if temporaryAgentError e then retryAction else atomically $ do - putTMVar clientVar (Left e) - -- TODO This can result in removing some other client from the map. - -- We need to identify these clients before they are connected and only remove if it's the same client in the map. - -- probably ClientVar needs it's own ID at a point it's created, and not rely on session ID of the connected client. - TM.delete tSess clients + putTMVar (sessionVar v) (Left e) + removeTSessVar v tSess clients throwError e tryConnectAsync :: m () tryConnectAsync = newAsyncAction connectAsync $ asyncClients c @@ -736,8 +724,8 @@ closeAgentClient c = liftIO $ do clearWorkers workers = atomically $ swapTVar (workers c) mempty clear :: Monoid m => (AgentClient -> TVar m) -> IO () clear sel = atomically $ writeTVar (sel c) mempty - cancelReconnect :: TMVar SubWorker -> IO () - cancelReconnect v = void . forkIO $ atomically (readTMVar v) >>= \(SubWorker _ a) -> uninterruptibleCancel a + cancelReconnect :: SessionVar (Async ()) -> IO () + cancelReconnect v = void . forkIO $ atomically (readTMVar $ sessionVar v) >>= uninterruptibleCancel cancelWorker :: Worker -> IO () cancelWorker Worker {doWork, action} = do @@ -765,9 +753,9 @@ closeClient c clientSel tSess = atomically (TM.lookupDelete tSess $ clientSel c) >>= mapM_ (closeClient_ c) closeClient_ :: ProtocolServerClient err msg => AgentClient -> ClientVar msg -> IO () -closeClient_ c cVar = do +closeClient_ c v = do NetworkConfig {tcpConnectTimeout} <- readTVarIO $ useNetworkConfig c - tcpConnectTimeout `timeout` atomically (readTMVar cVar) >>= \case + tcpConnectTimeout `timeout` atomically (readTMVar $ sessionVar v) >>= \case Just (Right client) -> closeProtocolServerClient client `catchAll_` pure () _ -> pure ()