diff --git a/src/Simplex/Messaging/Client.hs b/src/Simplex/Messaging/Client.hs index 454df91e8..f54fe686f 100644 --- a/src/Simplex/Messaging/Client.hs +++ b/src/Simplex/Messaging/Client.hs @@ -97,7 +97,7 @@ import Data.List (find) import Data.List.NonEmpty (NonEmpty (..)) import qualified Data.List.NonEmpty as L import Data.Maybe (fromMaybe) -import Data.Time.Clock (UTCTime (..), getCurrentTime) +import Data.Time.Clock (UTCTime (..), diffUTCTime, getCurrentTime) import Network.Socket (ServiceName) import Numeric.Natural import qualified Simplex.Messaging.Crypto as C @@ -110,7 +110,7 @@ import Simplex.Messaging.Transport import Simplex.Messaging.Transport.Client (SocksProxy, TransportClientConfig (..), TransportHost (..), defaultTcpConnectTimeout, runTransportClient) import Simplex.Messaging.Transport.KeepAlive import Simplex.Messaging.Transport.WebSockets (WS) -import Simplex.Messaging.Util (bshow, raceAny_, threadDelay', whenM) +import Simplex.Messaging.Util (bshow, diffToMicroseconds, raceAny_, threadDelay', whenM) import Simplex.Messaging.Version import System.Timeout (timeout) import UnliftIO (pooledMapConcurrentlyN) @@ -131,7 +131,9 @@ data PClient v err msg = PClient transportHost :: TransportHost, tcpTimeout :: Int, rcvConcurrency :: Int, - pingErrorCount :: TVar Int, + sendPings :: TVar Bool, + lastReceived :: TVar UTCTime, + timeoutErrorCount :: TVar Int, clientCorrId :: TVar ChaChaDRG, sentCommands :: TMap CorrId (Request err msg), sndQ :: TBQueue (TVar Bool, ByteString), @@ -141,10 +143,13 @@ data PClient v err msg = PClient smpClientStub :: TVar ChaChaDRG -> ByteString -> VersionSMP -> Maybe (THandleAuth 'TClient) -> STM SMPClient smpClientStub g sessionId thVersion thAuth = do + let ts = UTCTime (read "2024-03-31") 0 connected <- newTVar False clientCorrId <- C.newRandomDRG g sentCommands <- TM.empty - pingErrorCount <- newTVar 0 + sendPings <- newTVar False + lastReceived <- newTVar ts + timeoutErrorCount <- newTVar 0 sndQ <- newTBQueue 100 rcvQ <- newTBQueue 100 return @@ -159,7 +164,7 @@ smpClientStub g sessionId thVersion thAuth = do implySessId = thVersion >= authCmdsSMPVersion, batch = True }, - sessionTs = UTCTime (read "2024-03-31") 0, + sessionTs = ts, client_ = PClient { connected, @@ -167,7 +172,9 @@ smpClientStub g sessionId thVersion thAuth = do transportHost = "localhost", tcpTimeout = 15_000_000, rcvConcurrency = 8, - pingErrorCount, + sendPings, + lastReceived, + timeoutErrorCount, clientCorrId, sentCommands, sndQ, @@ -215,7 +222,7 @@ data NetworkConfig = NetworkConfig tcpKeepAlive :: Maybe KeepAliveOpts, -- | period for SMP ping commands (microseconds, 0 to disable) smpPingInterval :: Int64, - -- | the count of PING errors after which SMP client terminates (and will be reconnected), 0 to disable + -- | the count of timeout errors after which SMP client terminates (and will be reconnected), 0 to disable smpPingCount :: Int, logTLSErrors :: Bool } @@ -329,15 +336,17 @@ getProtocolClient :: forall v err msg. Protocol v err msg => TVar ChaChaDRG -> T getProtocolClient g transportSession@(_, srv, _) cfg@ProtocolClientConfig {qSize, networkConfig, serverVRange, agreeSecret} msgQ disconnected = do case chooseTransportHost networkConfig (host srv) of Right useHost -> - (atomically (mkProtocolClient useHost) >>= runClient useTransport useHost) + (getCurrentTime >>= atomically . mkProtocolClient useHost >>= runClient useTransport useHost) `catch` \(e :: IOException) -> pure . Left $ PCEIOError e Left e -> pure $ Left e where NetworkConfig {tcpConnectTimeout, tcpTimeout, rcvConcurrency, smpPingInterval} = networkConfig - mkProtocolClient :: TransportHost -> STM (PClient v err msg) - mkProtocolClient transportHost = do + mkProtocolClient :: TransportHost -> UTCTime -> STM (PClient v err msg) + mkProtocolClient transportHost ts = do connected <- newTVar False - pingErrorCount <- newTVar 0 + sendPings <- newTVar False + lastReceived <- newTVar ts + timeoutErrorCount <- newTVar 0 clientCorrId <- C.newRandomDRG g sentCommands <- TM.empty sndQ <- newTBQueue qSize @@ -348,7 +357,9 @@ getProtocolClient g transportSession@(_, srv, _) cfg@ProtocolClientConfig {qSize transportSession, transportHost, tcpTimeout, - pingErrorCount, + sendPings, + lastReceived, + timeoutErrorCount, clientCorrId, sentCommands, rcvConcurrency, @@ -386,6 +397,7 @@ getProtocolClient g transportSession@(_, srv, _) cfg@ProtocolClientConfig {qSize Right th@THandle {params} -> do sessionTs <- getCurrentTime let c' = ProtocolClient {action = Nothing, client_ = c, thParams = params, sessionTs} + atomically $ writeTVar (lastReceived c) sessionTs atomically $ do writeTVar (connected c) True putTMVar cVar $ Right c' @@ -396,17 +408,29 @@ getProtocolClient g transportSession@(_, srv, _) cfg@ProtocolClientConfig {qSize send ProtocolClient {client_ = PClient {sndQ}} h = forever $ atomically (readTBQueue sndQ) >>= \(active, s) -> whenM (readTVarIO active) (void $ tPutLog h s) receive :: Transport c => ProtocolClient v err msg -> THandle v c 'TClient -> IO () - receive ProtocolClient {client_ = PClient {rcvQ}} h = forever $ tGet h >>= atomically . writeTBQueue rcvQ + receive ProtocolClient {client_ = PClient {rcvQ, lastReceived, timeoutErrorCount}} h = forever $ do + tGet h >>= atomically . writeTBQueue rcvQ + getCurrentTime >>= atomically . writeTVar lastReceived + atomically $ writeTVar timeoutErrorCount 0 ping :: ProtocolClient v err msg -> IO () - ping c@ProtocolClient {client_ = PClient {pingErrorCount}} = do - threadDelay' smpPingInterval - runExceptT (sendProtocolCommand c Nothing "" $ protocolPing @v @err @msg) >>= \case - Left PCEResponseTimeout -> do - cnt <- atomically $ stateTVar pingErrorCount $ \cnt -> (cnt + 1, cnt + 1) - when (maxCnt == 0 || cnt < maxCnt) $ ping c - _ -> ping c -- sendProtocolCommand resets pingErrorCount + ping c@ProtocolClient {client_ = PClient {sendPings, lastReceived, timeoutErrorCount}} = loop smpPingInterval where + loop :: Int64 -> IO () + loop delay = do + threadDelay' delay + diff <- diffUTCTime <$> getCurrentTime <*> readTVarIO lastReceived + let idle = diffToMicroseconds diff + remaining = smpPingInterval - idle + if remaining > 1_000_000 -- delay pings only for significant time + then loop remaining + else do + whenM (readTVarIO sendPings) $ void . runExceptT $ sendProtocolCommand c Nothing "" (protocolPing @v @err @msg) + -- sendProtocolCommand/getResponse updates counter for each command + cnt <- readTVarIO timeoutErrorCount + -- drop client when maxCnt of commands have timed out in sequence, but only after some time has passed after last received response + when (maxCnt == 0 || cnt < maxCnt || diff < recoverWindow) $ loop smpPingInterval + recoverWindow = 15 * 60 -- seconds maxCnt = smpPingCount networkConfig process :: ProtocolClient v err msg -> IO () @@ -504,15 +528,18 @@ createSMPQueue c (rKey, rpKey) dhKey auth subMode = -- -- https://github.com/simplex-chat/simplexmq/blob/master/protocol/simplex-messaging.md#subscribe-to-queue subscribeSMPQueue :: SMPClient -> RcvPrivateAuthKey -> RecipientId -> ExceptT SMPClientError IO () -subscribeSMPQueue c rpKey rId = +subscribeSMPQueue c@ProtocolClient {client_ = PClient {sendPings}} rpKey rId = do + liftIO . atomically $ writeTVar sendPings True sendSMPCommand c (Just rpKey) rId SUB >>= \case - OK -> return () + OK -> pure () cmd@MSG {} -> liftIO $ writeSMPMessage c rId cmd r -> throwE . PCEUnexpectedResponse $ bshow r -- | Subscribe to multiple SMP queues batching commands if supported. subscribeSMPQueues :: SMPClient -> NonEmpty (RcvPrivateAuthKey, RecipientId) -> IO (NonEmpty (Either SMPClientError ())) -subscribeSMPQueues c qs = sendProtocolCommands c cs >>= mapM (processSUBResponse c) +subscribeSMPQueues c@ProtocolClient {client_ = PClient {sendPings}} qs = do + atomically $ writeTVar sendPings True + sendProtocolCommands c cs >>= mapM (processSUBResponse c) where cs = L.map (\(rpKey, rId) -> (Just rpKey, rId, Cmd SRecipient SUB)) qs @@ -720,11 +747,14 @@ sendProtocolCommand c@ProtocolClient {client_ = PClient {sndQ}, thParams = THand -- TODO switch to timeout or TimeManager that supports Int64 getResponse :: ProtocolClient v err msg -> TVar Bool -> Request err msg -> IO (Response err msg) -getResponse ProtocolClient {client_ = PClient {tcpTimeout, pingErrorCount, sentCommands}} active Request {corrId, entityId, responseVar} = do +getResponse ProtocolClient {client_ = PClient {tcpTimeout, timeoutErrorCount, sentCommands}} active Request {corrId, entityId, responseVar} = do response <- timeout tcpTimeout (atomically (takeTMVar responseVar)) >>= \case - Just r -> atomically (writeTVar pingErrorCount 0) $> r - Nothing -> atomically (writeTVar active False >> TM.delete corrId sentCommands) $> Left PCEResponseTimeout + Just r -> atomically (writeTVar timeoutErrorCount 0) $> r + Nothing -> do + atomically (writeTVar active False >> TM.delete corrId sentCommands) + atomically $ modifyTVar' timeoutErrorCount (+ 1) + pure $ Left PCEResponseTimeout pure Response {entityId, response} mkTransmission :: forall v err msg. ProtocolEncoding v err (ProtoCommand msg) => ProtocolClient v err msg -> ClientCommand msg -> IO (PCTransmission err msg)