cancel subscribers when client disconnects, reject ACK if MSG was not delivered

This commit is contained in:
Evgeny Poberezkin
2020-10-21 10:08:50 +01:00
parent ff49009be1
commit 2527cf8a65
4 changed files with 106 additions and 47 deletions
+2 -1
View File
@@ -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
+18 -6
View File
@@ -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
+79 -38
View File
@@ -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)
+7 -2
View File
@@ -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