SPDX-License-Identifier: AGPL-3.0-only

Give AgentClient a completion signal and make its shutdown join reconnect
workers, protocol clients and each concrete worker execution.  A per-worker
cancellation flag makes cancellation idempotent while allowing its first caller
to wait for a concurrent startup to publish the thread it must stop.  Worker
startup and late client publication are made atomic with shutdown, preventing
detached TLS and store users from outliving their owner.

Index: src/Simplex/Messaging/Agent/Client.hs
--- src/Simplex/Messaging/Agent/Client.hs.orig
+++ src/Simplex/Messaging/Agent/Client.hs
@@ -195,7 +195,7 @@
 import Control.Applicative ((<|>))
 import Control.Concurrent (ThreadId, killThread)
 import Control.Concurrent.Async (Async, uninterruptibleCancel)
-import Control.Concurrent.STM (retry)
+import Control.Concurrent.STM (retry, throwSTM)
 import Control.Exception (AsyncException (..), BlockedIndefinitelyOnSTM (..))
 import Control.Logger.Simple
 import Control.Monad
@@ -317,7 +317,7 @@
 import System.Random (randomR)
 import UnliftIO (mapConcurrently, timeout)
 import UnliftIO.Async (async)
-import UnliftIO.Concurrent (forkIO, mkWeakThreadId)
+import UnliftIO.Concurrent (forkFinally, forkIO, mkWeakThreadId)
 import UnliftIO.Directory (doesFileExist, getTemporaryDirectory, removeFile)
 import qualified UnliftIO.Exception as E
 import UnliftIO.STM
@@ -339,6 +339,7 @@

 data AgentClient = AgentClient
   { acThread :: TVar (Maybe (Weak ThreadId)),
+    acThreadDone :: TMVar (),
     active :: TVar Bool,
     subQ :: TBQueue ATransmission,
     msgQ :: TBQueue (ServerTransmissionBatch SMPVersion ErrorType BrokerMsg),
@@ -402,8 +403,8 @@
 {-# INLINE getAgentWorker #-}

 getAgentWorker' :: forall a k e m. (Ord k, Show k, AnyError e, MonadUnliftIO m) => (a -> Worker) -> (Worker -> STM a) -> String -> Bool -> AgentClient -> k -> TMap k a -> (a -> ExceptT e m ()) -> m a
-getAgentWorker' toW fromW name hasWork c@AgentClient {agentEnv} key ws work = do
-  atomically (getWorker >>= maybe createWorker whenExists) >>= \w -> runWorker w $> w
+getAgentWorker' toW fromW name hasWork c@AgentClient {active, agentEnv} key ws work =
+  E.mask_ $ atomically (unlessM (readTVar active) (throwSTM ThreadKilled) >> getWorker >>= maybe createWorker whenExists) >>= \w -> runWorker w $> w
   where
     getWorker = TM.lookup key ws
     createWorker = do
@@ -413,7 +414,7 @@
     whenExists w
       | hasWork = hasWorkToDo (toW w) $> w
       | otherwise = pure w
-    runWorker w = runWorkerAsync (toW w) runWork
+    runWorker w = runWorkerAsync active (toW w) runWork
       where
         runWork :: m ()
         runWork = tryAllErrors' (work w) >>= restartOrDelete
@@ -428,6 +428,6 @@
           | wId == workerId (toW w') = do
               rc <- readTVar restarts
-              isActive <- readTVar $ active c
+              isActive <- readTVar active
               checkRestarts isActive $ updateRestartCount t rc
           | otherwise =
               pure False -- there is a new worker in the map, no action
@@ -456,15 +457,18 @@
-  action <- newTMVar Nothing
+  (action, cancelled) <- (,) <$> newTMVar Nothing <*> newTVar False
   restarts <- newTVar $ RestartCount 0 0
-  pure Worker {workerId, doWork, action, restarts}
+  pure Worker {workerId, doWork, action, cancelled, restarts}
 
-runWorkerAsync :: MonadUnliftIO m => Worker -> m () -> m ()
-runWorkerAsync Worker {action} work =
+runWorkerAsync :: MonadUnliftIO m => TVar Bool -> Worker -> m () -> m ()
+runWorkerAsync active Worker {action, cancelled} work =
   E.bracket
-    (atomically $ takeTMVar action) -- get current action, locking to avoid race conditions
+    (atomically $ unlessM (readTVar active) (throwSTM ThreadKilled) >> whenM (readTVar cancelled) (throwSTM ThreadKilled) >> takeTMVar action) -- get current action, locking to avoid race conditions
     (atomically . tryPutTMVar action) -- if it was running (or if start crashes), put it back and unlock (don't lock if it was just started)
     (\a -> when (isNothing a) start) -- start worker if it's not running
   where
-    start = atomically . putTMVar action . Just =<< mkWeakThreadId =<< forkIO work
+    start = E.mask_ $ do
+      done <- newEmptyTMVarIO
+      t <- work `forkFinally` const (atomically $ putTMVar done ())
+      atomically . putTMVar action . Just . (,done) =<< mkWeakThreadId t
 
 data AgentOperation = AONtfNetwork | AORcvNetwork | AOMsgDelivery | AOSndNetwork | AODatabase
   deriving (Eq, Show)
@@ -514,6 +518,7 @@
       qSize = tbqSize cfg
   proxySessTs <- newTVarIO =<< getCurrentTime
   acThread <- newTVarIO Nothing
+  acThreadDone <- newEmptyTMVarIO
   active <- newTVarIO True
   subQ <- newTBQueueIO qSize
   msgQ <- newTBQueueIO qSize
@@ -554,6 +559,7 @@
   return
     AgentClient
       { acThread,
+        acThreadDone,
         active,
         subQ,
         msgQ,
@@ -798,7 +804,7 @@
 resubscribeSMPSession :: AgentClient -> SMPTransportSession -> AM' ()
 resubscribeSMPSession c@AgentClient {smpSubWorkers, workerSeq} tSess = do
   ts <- liftIO getCurrentTime
-  atomically (getWorkerVar ts) >>= mapM_ (either newSubWorker (\_ -> pure ()))
+  E.mask_ $ atomically (ifM (readTVar $ active c) (getWorkerVar ts) (pure Nothing)) >>= mapM_ (either newSubWorker (\_ -> pure ()))
   where
     getWorkerVar ts =
       ifM
@@ -858,7 +864,7 @@
     clientDisconnected :: NtfClientVar -> NtfClient -> IO ()
     clientDisconnected v client = do
       atomically $ removeSessVar v tSess ntfClients
-      atomically $ writeTBQueue (subQ c) ("", "", AEvt SAENone $ hostEvent DISCONNECT client)
+      whenM (readTVarIO active) $ atomically $ writeTBQueue (subQ c) ("", "", AEvt SAENone $ hostEvent DISCONNECT client)
       logInfo . decodeUtf8 $ "Agent disconnected from " <> showServer srv
 
 getXFTPServerClient :: AgentClient -> XFTPTransportSession -> AM XFTPClient
@@ -879,7 +885,7 @@
     clientDisconnected :: XFTPClientVar -> XFTPClient -> IO ()
     clientDisconnected v client = do
       atomically $ removeSessVar v tSess xftpClients
-      atomically $ writeTBQueue (subQ c) ("", "", AEvt SAENone $ hostEvent DISCONNECT client)
+      whenM (readTVarIO active) $ atomically $ writeTBQueue (subQ c) ("", "", AEvt SAENone $ hostEvent DISCONNECT client)
       logInfo . decodeUtf8 $ "Agent disconnected from " <> showServer srv
 
 waitForProtocolClient ::
@@ -916,10 +922,16 @@
 newProtocolClient c tSess@(userId, srv, entityId_) clients connectClient v =
   tryAllErrors (connectClient v) >>= \case
     Right client -> do
-      logInfo . decodeUtf8 $ "Agent connected to " <> showServer srv <> " (user " <> bshow userId <> maybe "" (" for entity " <>) entityId_ <> ")"
-      atomically $ putTMVar (sessionVar v) (Right client)
-      liftIO $ nonBlockingWriteTBQueue (subQ c) ("", "", AEvt SAENone $ hostEvent CONNECT client)
-      pure client
+      isActive <- atomically $ do
+        isActive <- readTVar $ active c
+        when isActive $ putTMVar (sessionVar v) (Right client)
+        pure isActive
+      if isActive
+        then do
+          logInfo . decodeUtf8 $ "Agent connected to " <> showServer srv <> " (user " <> bshow userId <> maybe "" (" for entity " <>) entityId_ <> ")"
+          liftIO $ nonBlockingWriteTBQueue (subQ c) ("", "", AEvt SAENone $ hostEvent CONNECT client)
+          pure client
+        else liftIO $ closeProtocolServerClient (protocolClient client) `catchAll_` pure () >> E.throwIO ThreadKilled
     Left e -> do
       ei <- asks $ persistErrorInterval . config
       if ei == 0
@@ -967,11 +979,11 @@
 closeAgentClient :: AgentClient -> IO ()
 closeAgentClient c = do
   atomically $ writeTVar (active c) False
+  atomically (swapTVar (smpSubWorkers c) M.empty) >>= void . mapConcurrently cancelReconnect
   closeProtocolServerClients c smpClients
   closeProtocolServerClients c ntfClients
   closeProtocolServerClients c xftpClients
   atomically $ writeTVar (smpProxiedRelays c) M.empty
-  atomically (swapTVar (smpSubWorkers c) M.empty) >>= mapM_ cancelReconnect
   clearWorkers smpDeliveryWorkers >>= mapM_ (cancelWorker . fst)
   clearWorkers asyncCmdWorkers >>= mapM_ cancelWorker
   atomically $ SS.clear $ currentSubs c
@@ -983,12 +995,12 @@
     clear :: Monoid m => (AgentClient -> TVar m) -> IO ()
     clear sel = atomically $ writeTVar (sel c) mempty
     cancelReconnect :: SessionVar (Async ()) -> IO ()
-    cancelReconnect v = void . forkIO $ atomically (readTMVar $ sessionVar v) >>= uninterruptibleCancel
+    cancelReconnect v = atomically (readTMVar $ sessionVar v) >>= uninterruptibleCancel
 
 cancelWorker :: Worker -> IO ()
-cancelWorker Worker {doWork, action} = do
+cancelWorker Worker {doWork, action, cancelled} = do
   noWorkToDo doWork
-  atomically (tryTakeTMVar action) >>= mapM_ (mapM_ $ deRefWeak >=> mapM_ killThread)
+  atomically (ifM (stateTVar cancelled $ \b -> (b, True)) (pure Nothing) (Just <$> takeTMVar action)) >>= mapM_ (mapM_ $ \(t, done) -> deRefWeak t >>= mapM_ killThread >> atomically (readTMVar done))
 
 waitUntilActive :: AgentClient -> IO ()
 waitUntilActive AgentClient {active} = unlessM (readTVarIO active) $ atomically $ unlessM (readTVar active) retry
@@ -1005,7 +1017,7 @@

 closeProtocolServerClients :: ProtocolServerClient v err msg => AgentClient -> (AgentClient -> TMap (TransportSession msg) (ClientVar msg)) -> IO ()
 closeProtocolServerClients c clientsSel =
-  atomically (clientsSel c `swapTVar` M.empty) >>= mapM_ (forkIO . closeClient_ c)
+  atomically (clientsSel c `swapTVar` M.empty) >>= void . mapConcurrently (closeClient_ c)
 
 reconnectServerClients :: ProtocolServerClient v err msg => AgentClient -> (AgentClient -> TMap (TransportSession msg) (ClientVar msg)) -> IO ()
 reconnectServerClients c clientsSel =
