diff --git a/src/Simplex/Messaging/Server.hs b/src/Simplex/Messaging/Server.hs index 90b2c822c..85a262717 100644 --- a/src/Simplex/Messaging/Server.hs +++ b/src/Simplex/Messaging/Server.hs @@ -198,6 +198,7 @@ smpServer started cfg@ServerConfig {transports, transportConfig = tCfg, startOpt : sigIntHandlerThread : map runServer transports <> expireMessagesThread_ cfg + <> expireClientsThread_ cfg <> serverStatsThread_ cfg <> prometheusMetricsThread_ cfg <> controlPortThread_ cfg @@ -475,6 +476,27 @@ smpServer started cfg@ServerConfig {transports, transportConfig = tCfg, startOpt expireMessagesThread_ ServerConfig {messageExpiration = Just msgExp} = [expireMessagesThread msgExp] expireMessagesThread_ _ = [] + expireClientsThread_ :: ServerConfig s -> [M s ()] + expireClientsThread_ ServerConfig {inactiveClientExpiration = Just expCfg} = [expireClientsThread expCfg] + expireClientsThread_ _ = [] + + expireClientsThread :: ExpirationConfig -> M s () + expireClientsThread expCfg = do + labelMyThread "expireClients" + srv <- asks server + liftIO $ forever $ do + threadDelay' $ checkInterval expCfg * 1000000 + old <- expireBeforeEpoch expCfg + getServerClients srv >>= mapM_ (expireClient srv old) + where + expireClient srv old c@Client {rcvActiveAt, sndActiveAt, closeTransport} = do + ts <- max <$> readTVarIO rcvActiveAt <*> readTVarIO sndActiveAt + when (systemSeconds ts < old) $ whenM (noSubscriptions srv c) closeTransport + noSubscriptions srv Client {clientId} = + not <$> anyM [hasSubs (subscribers srv), hasSubs (ntfSubscribers srv)] + where + hasSubs ServerSubscribers {subClients} = IS.member clientId <$> readTVarIO subClients + expireMessagesThread :: ExpirationConfig -> M s () expireMessagesThread ExpirationConfig {checkInterval, ttl} = do ms <- asks msgStore @@ -1068,7 +1090,7 @@ runClientTransport h@THandle {params = thParams@THandleParams {sessionId}} = do ts <- liftIO getSystemTime nextClientId <- asks clientSeq clientId <- atomically $ stateTVar nextClientId $ \next -> (next, next + 1) - c <- liftIO $ newClient clientId q thParams ts + c <- liftIO $ newClient clientId q thParams ts (closeConnection $ connection h) runClientThreads c `finally` clientDisconnected c where runClientThreads :: Client s -> M s () @@ -1076,16 +1098,8 @@ runClientTransport h@THandle {params = thParams@THandleParams {sessionId}} = do s <- asks server ms <- asks msgStore whenM (liftIO $ insertServerClient c s) $ do - expCfg <- asks $ inactiveClientExpiration . config labelMyThread . B.unpack $ "client $" <> encode sessionId - raceAny_ $ [liftIO $ send h c, client s ms c, receive h ms c] <> disconnectThread_ c s expCfg - disconnectThread_ :: Client s -> Server s -> Maybe ExpirationConfig -> [M s ()] - disconnectThread_ c s (Just expCfg) = [liftIO $ disconnectTransport h (rcvActiveAt c) (sndActiveAt c) expCfg (noSubscriptions c s)] - disconnectThread_ _ _ _ = [] - noSubscriptions Client {clientId} s = - not <$> anyM [hasSubs (subscribers s), hasSubs (ntfSubscribers s)] - where - hasSubs ServerSubscribers {subClients} = IS.member clientId <$> readTVarIO subClients + raceAny_ [liftIO $ send h c, client s ms c, receive h ms c] controlPortAuth :: Handle -> Maybe BasicAuth -> Maybe BasicAuth -> TVar CPClientRole -> BasicAuth -> IO () controlPortAuth h user admin role auth = do diff --git a/src/Simplex/Messaging/Server/Env/STM.hs b/src/Simplex/Messaging/Server/Env/STM.hs index b4333959f..52f4c45f6 100644 --- a/src/Simplex/Messaging/Server/Env/STM.hs +++ b/src/Simplex/Messaging/Server/Env/STM.hs @@ -465,7 +465,8 @@ data Client s = Client connected :: TVar Bool, createdAt :: SystemTime, rcvActiveAt :: TVar SystemTime, - sndActiveAt :: TVar SystemTime + sndActiveAt :: TVar SystemTime, + closeTransport :: IO () } type VerifiedTransmission s = (Maybe (StoreQueue s, QueueRec), Transmission Cmd) @@ -520,8 +521,8 @@ newServerSubscribers = do pendingEvents <- newTVarIO IM.empty pure ServerSubscribers {subQ, queueSubscribers, serviceSubscribers, totalServiceSubs, subClients, pendingEvents} -newClient :: ClientId -> Natural -> THandleParams SMPVersion 'TServer -> SystemTime -> IO (Client s) -newClient clientId qSize clientTHParams createdAt = do +newClient :: ClientId -> Natural -> THandleParams SMPVersion 'TServer -> SystemTime -> IO () -> IO (Client s) +newClient clientId qSize clientTHParams createdAt closeTransport = do subscriptions <- TM.emptyIO ntfSubscriptions <- TM.emptyIO serviceSubscribed <- newTVarIO False @@ -556,7 +557,8 @@ newClient clientId qSize clientTHParams createdAt = do connected, createdAt, rcvActiveAt, - sndActiveAt + sndActiveAt, + closeTransport } newSubscription :: SubscriptionThread -> STM Sub diff --git a/tests/ServerTests.hs b/tests/ServerTests.hs index 116b4f0ec..3930056c6 100644 --- a/tests/ServerTests.hs +++ b/tests/ServerTests.hs @@ -106,6 +106,7 @@ serverTests = do testMsgExpireOnInterval testMsgNOTExpireOnInterval describe "Blocking queues" $ testBlockMessageQueue + describe "Inactive clients" testInactiveClientExpiration describe "Short links" $ do testInvQueueLinkData testContactQueueLinkData @@ -1563,6 +1564,23 @@ testMsgExpireOnInterval = Nothing -> return () Just _ -> error "nothing should be delivered" +testInactiveClientExpiration :: SpecWith (ASrvTransport, AStoreType) +testInactiveClientExpiration = + it "should disconnect inactive clients without subscriptions" $ \(ATransport (t :: TProxy c 'TServer), msType) -> do + g <- C.newRandom + (rPub, rKey) <- atomically $ C.generateAuthKeyPair C.SEd25519 g + (dhPub, _ :: C.PrivateKeyX25519) <- atomically $ C.generateKeyPair g + let cfg' = updateCfg (cfgMS msType) $ \cfg_ -> cfg_ {inactiveClientExpiration = Just ExpirationConfig {ttl = 1, checkInterval = 1}} + withSmpServerConfigOn (ATransport t) cfg' testPort $ \_ -> + testSMPClient @c $ \rh -> testSMPClient @c $ \h -> do + Resp "1" NoEntity (Ids _ _ _) <- signSendRecv rh rKey ("1", NoEntity, New rPub dhPub) + threadDelay 2500000 + try (timeout 2000000 $ tGet1 h) >>= \case + Left (_ :: SomeException) -> pure () + Right r -> unexpected (r :: Maybe (Transmission (Either ErrorType BrokerMsg))) + Resp "2" NoEntity PONG <- sendRecv rh (Nothing, "2", NoEntity, PING) + pure () + testMsgNOTExpireOnInterval :: SpecWith (ASrvTransport, AStoreType) testMsgNOTExpireOnInterval = it "should block and unblock message queues" $ \(ATransport (t :: TProxy c 'TServer), msType) -> do