diff --git a/src/Simplex/Messaging/Server.hs b/src/Simplex/Messaging/Server.hs index 50a9c5110..ee256da3f 100644 --- a/src/Simplex/Messaging/Server.hs +++ b/src/Simplex/Messaging/Server.hs @@ -1476,11 +1476,15 @@ client -- Run a slow command on a thread forkCmd :: (ServerConfig s -> Int) -> CorrId -> EntityId -> M s BrokerMsg -> M s (Maybe a) forkCmd concurrency corrId entId cmdAction = do - bracket_ wait signal . forkClient clnt (B.unpack $ "client $" <> encode sessionId <> " cmd") $ - -- commands MUST be processed under a reasonable timeout or the client would halt - cmdAction >>= \t -> atomically $ writeTBQueue sndQ ([(corrId, entId, t)], []) + -- the forked thread releases the slot when the command completes, the caller only if the fork failed + mask $ \restore -> do + wait + forkClient clnt (B.unpack $ "client $" <> encode sessionId <> " cmd") (restore cmd `finally` signal) + `onException` signal pure Nothing where + -- commands MUST be processed under a reasonable timeout or the client would halt + cmd = cmdAction >>= \t -> atomically $ writeTBQueue sndQ ([(corrId, entId, t)], []) wait = do limit <- asks (concurrency . config) atomically $ do diff --git a/tests/RSLVTests.hs b/tests/RSLVTests.hs index 9263b6953..96bd0639f 100644 --- a/tests/RSLVTests.hs +++ b/tests/RSLVTests.hs @@ -11,11 +11,14 @@ module RSLVTests (rslvTests) where +import Control.Concurrent (threadDelay) +import Control.Monad (forM_) import Control.Monad.Trans.Except (ExceptT, runExceptT) import qualified Data.ByteString.Char8 as B import qualified Data.ByteString.Lazy as LB import Data.IORef (IORef, readIORef) import Data.List.NonEmpty (NonEmpty (..)) +import qualified Data.List.NonEmpty as L import Data.Text (Text) import Data.Text.Encoding (encodeUtf8) import Data.Time.Clock (getCurrentTime) @@ -47,6 +50,7 @@ import Simplex.Messaging.Protocol tPut, ) import qualified Simplex.Messaging.Protocol as SMP +import Simplex.Messaging.Server.Env.STM (ServerConfig (..)) import Simplex.Messaging.SimplexName (SimplexDomain) import Simplex.Messaging.Transport import Simplex.Messaging.Version (mkVersionRange) @@ -102,6 +106,8 @@ rslvTests = do it "RSLV sends the 2LD as its hash" testRslvSendsTheHash it "a name with subnames is sent as text" testSubnameKeepsItsLabels it "a record naming a different name is rejected" testRslvWrongName + describe "RSLV resource use" $ + it "one connection has at most resolver_concurrency lookups in flight" testRslvConnectionCap -- | /v2/resolve answers 200, 400 or 502, so a 404 is a resolver that predates -- the route, not a name that does not exist. @@ -275,5 +281,34 @@ testRslvWrongName = Left (PCEUnexpectedResponse _) -> pure () _ -> expectationFailure $ "expected Left (PCEUnexpectedResponse ..), got: " <> show r +-- The resolver answers after 3s and the server gives up after 1s, so a request +-- that reached the resolver stays in flight for the whole check. +testRslvConnectionCap :: IO () +testRslvConnectionCap = + NRS.withResolverServerDelayed 3000 (NRS.resolveResp status200 "{}") $ \port reqs -> + withSmpServerConfigOn (transport @TLS) (updateCfg (withNames port memCfg) $ \c -> c {serverResolverConcurrency = connCap}) testPort $ const $ + testSMPClient @TLS $ \h -> do + sendRslvs h "cap" 16 + threadDelay 800000 + length <$> resolvePaths reqs `shouldReturn` connCap + recvResponses h 16 `shouldReturn` replicate 16 (Right (ERR (NAME (RESOLVER "timeout")))) + where + connCap = 4 + +-- | One RSLV per block, so no batch limit applies. +sendRslvs :: THandleSMP TLS 'TClient -> String -> Int -> IO () +sendRslvs h@THandle {params} prefix n = + forM_ [1 .. n] $ \i -> do + let TransmissionForAuth {tToSend} = encodeTransmissionForAuth params (CorrId (B.pack $ prefix <> show i), NoEntity, Cmd SResolver (RSLV (NQDomain (domain "alice.simplex")))) + [Right ()] <- tPut h (Right (Nothing, tToSend) :| []) + pure () + +recvResponses :: THandleSMP TLS 'TClient -> Int -> IO [Either ErrorType BrokerMsg] +recvResponses h n + | n <= 0 = pure [] + | otherwise = do + rs <- map (\(_, _, r) -> r) . L.toList <$> tGetClient h + (rs <>) <$> recvResponses h (n - length rs) + runExceptT' :: Show e => ExceptT e IO a -> IO a runExceptT' a = runExceptT a >>= either (fail . show) pure