diff --git a/package.yaml b/package.yaml index b6333b893..9c0bf214a 100644 --- a/package.yaml +++ b/package.yaml @@ -29,6 +29,8 @@ executables: simplex-messaging: source-dirs: src main: Main.hs + ghc-options: + - -threaded library: source-dirs: src @@ -45,7 +47,6 @@ tests: ghc-options: # - -haddock - - -O0 - -Wall - -Wcompat - -Werror=incomplete-patterns diff --git a/src/Env/STM.hs b/src/Env/STM.hs index f651ca295..378df6232 100644 --- a/src/Env/STM.hs +++ b/src/Env/STM.hs @@ -32,27 +32,39 @@ data Env = Env data Server = Server { subscribedQ :: TBQueue (RecipientId, Client), - connections :: TVar (Map RecipientId Client) + subscribers :: TVar (Map RecipientId Client) } data Client = Client - { connections :: TVar (Map RecipientId (Either () ThreadId)), + { subscriptions :: TVar (Map RecipientId Sub), rcvQ :: TBQueue Signed, sndQ :: TBQueue Signed } +data SubscriptionThread = NoSub | SubPending | SubThread ThreadId + +data Sub = Sub + { subThread :: SubscriptionThread, + delivered :: TMVar () + } + newServer :: Natural -> STM Server newServer qSize = do subscribedQ <- newTBQueue qSize - connections <- newTVar M.empty - return Server {subscribedQ, connections} + subscribers <- newTVar M.empty + return Server {subscribedQ, subscribers} newClient :: Natural -> STM Client newClient qSize = do - connections <- newTVar M.empty + subscriptions <- newTVar M.empty rcvQ <- newTBQueue qSize sndQ <- newTBQueue qSize - return Client {connections, rcvQ, sndQ} + return Client {subscriptions, rcvQ, sndQ} + +newSubscription :: STM Sub +newSubscription = do + delivered <- newEmptyTMVar + return Sub {subThread = NoSub, delivered} newEnv :: (MonadUnliftIO m, MonadRandom m) => Config -> m Env newEnv config = do diff --git a/src/Server.hs b/src/Server.hs index 6af048122..862ca6dd4 100644 --- a/src/Server.hs +++ b/src/Server.hs @@ -13,6 +13,7 @@ module Server (runSMPServer) where import ConnStore +import Control.Concurrent.STM (stateTVar) import Control.Monad import Control.Monad.IO.Unlift import Control.Monad.Reader @@ -28,6 +29,7 @@ import Transmission import Transport import UnliftIO.Async import UnliftIO.Concurrent +import UnliftIO.Exception import UnliftIO.IO import UnliftIO.STM @@ -42,13 +44,13 @@ runSMPServer cfg@Config {tcpPort} = do race_ (runTCPServer tcpPort runClient) (serverThread s) serverThread :: MonadUnliftIO m => Server -> m () - serverThread Server {subscribedQ, connections} = forever . atomically $ do + serverThread Server {subscribedQ, subscribers} = forever . atomically $ do (rId, clnt) <- readTBQueue subscribedQ - cs <- readTVar connections + cs <- readTVar subscribers case M.lookup rId cs of Just Client {rcvQ} -> writeTBQueue rcvQ (rId, Cmd SBroker END) Nothing -> return () - writeTVar connections $ M.insert rId clnt cs + writeTVar subscribers $ M.insert rId clnt cs runClient :: (MonadUnliftIO m, MonadReader Env m) => Handle -> m () runClient h = do @@ -57,6 +59,17 @@ runClient h = do c <- atomically $ newClient q s <- asks server raceAny_ [send h c, client c s, receive h c] + `finally` cancelSubscribers c + +cancelSubscribers :: (MonadUnliftIO m) => Client -> m () +cancelSubscribers Client {subscriptions} = do + cs <- readTVarIO subscriptions + forM_ cs cancelSub + +cancelSub :: (MonadUnliftIO m) => Sub -> m () +cancelSub = \case + Sub {subThread = SubThread t} -> killThread t + _ -> return () raceAny_ :: MonadUnliftIO m => [m a] -> m () raceAny_ = r [] @@ -103,7 +116,7 @@ verifyTransmission signature connId cmd = do authErr = smpErr AUTH client :: forall m. (MonadUnliftIO m, MonadReader Env m) => Client -> Server -> m () -client clnt@Client {connections, rcvQ, sndQ} Server {subscribedQ} = +client clnt@Client {subscriptions, rcvQ, sndQ} Server {subscribedQ} = forever $ atomically (readTBQueue rcvQ) >>= processCommand @@ -113,13 +126,13 @@ client clnt@Client {connections, rcvQ, sndQ} Server {subscribedQ} = processCommand (connId, cmd) = do st <- asks connStore case cmd of - Cmd SBroker END -> unsubscribeConn connId >> return (connId, cmd) + Cmd SBroker END -> unsubscribeConn >> return (connId, cmd) Cmd SBroker _ -> return (connId, cmd) - Cmd SSender (SEND msgBody) -> sendMessage st connId msgBody + Cmd SSender (SEND msgBody) -> sendMessage st msgBody Cmd SRecipient command -> case command of CONN rKey -> createConn st rKey SUB -> subscribeConn connId - ACK -> deliverMessage tryDelPeekMsg connId -- TODO? sending ACK without message loses the message + ACK -> acknowledgeMsg KEY sKey -> okResponse <$> secureConn st connId sKey OFF -> okResponse <$> suspendConn st connId DEL -> okResponse <$> deleteConn st connId @@ -146,24 +159,36 @@ client clnt@Client {connections, rcvQ, sndQ} Server {subscribedQ} = subscribeConn :: RecipientId -> m Signed subscribeConn rId = do atomically $ do - cs <- readTVar connections - when (M.notMember rId cs) $ do - writeTBQueue subscribedQ (rId, clnt) - writeTVar connections $ M.insert rId (Left ()) cs + cs <- readTVar subscriptions + case M.lookup rId cs of + Just Sub {delivered} -> void $ tryTakeTMVar delivered + Nothing -> do + writeTBQueue subscribedQ (rId, clnt) + sub <- newSubscription + writeTVar subscriptions $ M.insert rId sub cs deliverMessage tryPeekMsg rId - unsubscribeConn :: RecipientId -> m () - unsubscribeConn rId = do - cs <- readTVarIO connections - atomically . writeTVar connections $ M.delete rId cs - case M.lookup rId cs of - Just (Right threadId) -> killThread threadId - _ -> return () + unsubscribeConn :: m () + unsubscribeConn = do + sub <- atomically . stateTVar subscriptions $ + \cs -> (M.lookup connId cs, M.delete connId cs) + mapM_ cancelSub sub - sendMessage :: MonadConnStore s m => s -> SenderId -> MsgBody -> m Signed - sendMessage st sId msgBody = - getConn st SSender sId - >>= fmap (mkSigned sId) . either (return . ERR) storeMessage + acknowledgeMsg :: m Signed + acknowledgeMsg = do + dlvrd <- atomically $ do + cs <- readTVar subscriptions + case M.lookup connId cs of + Just s -> tryTakeTMVar (delivered s) + Nothing -> return Nothing + case dlvrd of + Just () -> deliverMessage tryDelPeekMsg connId + Nothing -> return . mkSigned connId $ ERR PROHIBITED + + sendMessage :: MonadConnStore s m => s -> MsgBody -> m Signed + sendMessage st msgBody = + getConn st SSender connId + >>= fmap (mkSigned connId) . either (return . ERR) storeMessage where mkMessage :: m Message mkMessage = do @@ -185,24 +210,40 @@ client clnt@Client {connections, rcvQ, sndQ} Server {subscribedQ} = ms <- asks msgStore q <- getMsgQueue ms rId tryPeek q >>= \case - Just msg -> return $ msgResponse rId msg - Nothing -> forkSubscriber q rId - - forkSubscriber :: MsgQueue -> RecipientId -> m Signed - forkSubscriber q rId = do - cs <- readTVarIO connections - case M.lookup rId cs of - Just (Left ()) -> do - threadId <- forkIO subscriber - trackSubscriber $ Right threadId - return ok - _ -> return ok + Just msg -> do + atomically $ do + sub <- M.lookup rId <$> readTVar subscriptions + forM_ sub $ \Sub {delivered} -> tryPutTMVar delivered () + return $ msgResponse rId msg + Nothing -> forkSub q >> return ok where - trackSubscriber sThrd = atomically . modifyTVar connections $ M.insert rId sThrd - subscriber = do + forkSub :: MsgQueue -> m () + forkSub q = do + sub <- M.lookup rId <$> readTVarIO subscriptions + case sub of + Just Sub {subThread = NoSub} -> do + atomically . setSub $ \s -> s {subThread = SubPending} + t <- forkIO $ subscriber q + atomically . setSub $ \case + s@Sub {subThread = SubPending} -> s {subThread = SubThread t} + s -> s + _ -> return () + + setSub :: (Sub -> Sub) -> STM () + setSub f = modifyTVar subscriptions $ M.adjust f rId + + subscriber :: MsgQueue -> m () + subscriber q = do msg <- peekMsg q - atomically . writeTBQueue sndQ $ msgResponse rId msg - trackSubscriber $ Left () + atomically $ do + writeTBQueue sndQ $ msgResponse rId msg + -- setSub (\s -> s {subThread = NoSub}) + cs <- readTVar subscriptions + let sub = M.lookup rId cs + forM_ sub $ \s@Sub {delivered} -> do + void $ tryPutTMVar delivered () + let cs' = M.insert rId s {subThread = NoSub} cs + writeTVar subscriptions cs' mkSigned :: ConnId -> Command 'Broker -> Signed mkSigned cId command = (cId, Cmd SBroker command) diff --git a/tests/Test.hs b/tests/Test.hs index e23e3464a..899653cfc 100644 --- a/tests/Test.hs +++ b/tests/Test.hs @@ -222,11 +222,16 @@ testSwitchSub = Resp _ (MSG _ _ msg3) <- tGet fromServer rh2 (msg3, "test3") #== "delivered to the 2nd TCP connection" - Resp _ OK <- sendRecv rh1 ("1234", rId, "ACK") + + Resp _ err <- sendRecv rh1 ("1234", rId, "ACK") + (err, ERR PROHIBITED) #== "rejects ACK from the 1st TCP connection" + + Resp _ ok3 <- sendRecv rh2 ("1234", rId, "ACK") + (ok3, OK) #== "accepts ACK from the 2nd TCP connection" timeout 1000 (tGet fromServer rh1) >>= \case Nothing -> return () - Just _ -> error "nothing should be delivered to the 1st TCPconnection" + Just _ -> error "nothing else is delivered to the 1st TCPconnection" syntaxTests :: SpecWith () syntaxTests = do