diff --git a/packages/network-transport-quic/CHANGELOG.md b/packages/network-transport-quic/CHANGELOG.md index d4e96dc1..cf192361 100644 --- a/packages/network-transport-quic/CHANGELOG.md +++ b/packages/network-transport-quic/CHANGELOG.md @@ -1,3 +1,11 @@ +Unreleased Laurent P. René de Cotret 0.2.0 + +* All the logical connections between two endpoints are now carried by a single QUIC connection (one stream + each), rather than by one QUIC connection for each endpoint pairs. This has large performance implications: + for multiple logical connections between two endpoints, `network-transport-quic` throughput increases by 50% over + version 0.1.x, for a total of 3x throughput over `network-transport-quic`. +* Breaking change: A new `socketOptions` field to `QUICTransportConfig`, allowing the user to control the UDP socket + underlying a connection. 2026-04-21 Laurent P. René de Cotret 0.1.2 diff --git a/packages/network-transport-quic/README.md b/packages/network-transport-quic/README.md index dd4d03ba..a6ce705b 100644 --- a/packages/network-transport-quic/README.md +++ b/packages/network-transport-quic/README.md @@ -8,7 +8,7 @@ QUIC has many advantages over TCP, including: * Connection migration. Connections survive IP address changes, which is important when a device switches from e.g. WIFI to 5G; * Built-in encryption via TLS 1.3; -In benchmarks, `network-transport-quic` performs better than `network-transport-tcp` in dense network topologies. For example, if every `EndPoint` in your network connects to every other `EndPoint`, you might benefit greatly from switching to `network-transport-quic`! +In benchmarks, `network-transport-quic` performs better than `network-transport-tcp` in dense network topologies. For multiple logical connections between two endpoints, `network-transport-quic` can be 3x faster (in throughput) compared to `network-transport-tcp`. ## Usage example diff --git a/packages/network-transport-quic/bench/Bench.hs b/packages/network-transport-quic/bench/Bench.hs index 53e10d92..7b5aaf23 100644 --- a/packages/network-transport-quic/bench/Bench.hs +++ b/packages/network-transport-quic/bench/Bench.hs @@ -6,41 +6,42 @@ module Main where -import Control.Concurrent (forkIO) -import Control.Concurrent.Async (forConcurrently_) +import Control.Concurrent.Async (forConcurrently_, link, wait, withAsync) import Control.Concurrent.MVar (newEmptyMVar, putMVar, takeMVar) import Control.Exception (finally, throwIO) -import Control.Monad (forM_, replicateM, void, when) +import Control.Monad (forM_, forever, replicateM, void, when) import qualified Data.ByteString as BS -import Data.IORef ( - atomicModifyIORef', - newIORef, - ) +import Data.IORef + ( atomicModifyIORef', + newIORef, + ) import Data.List.NonEmpty (NonEmpty (..)) -import Network.Transport ( - Connection (send), - EndPoint (address, connect, receive), - Event (ConnectionOpened, Received), - Reliability (ReliableOrdered), - Transport (closeTransport, newEndPoint), - defaultConnectHints, - ) +import qualified Network.Socket as N +import Network.Transport + ( Connection (send), + EndPoint (address, connect, receive), + Event (ConnectionOpened, ErrorEvent, Received), + Reliability (ReliableOrdered), + Transport (closeTransport, newEndPoint), + defaultConnectHints, + ) import qualified Network.Transport.QUIC as QUIC import qualified Network.Transport.TCP as TCP import System.FilePath (()) -import Test.Tasty (TestTree) -import Test.Tasty.Bench (bench, bgroup, defaultMain, nfIO) +import System.Timeout (timeout) +import Test.Tasty (localOption) +import Test.Tasty.Bench (Benchmark, TimeMode (WallTime), bench, bgroup, defaultMain, nfIO) data TransportConfig = TransportConfig - { transportName :: String - , mkTransport :: IO Transport + { transportName :: String, + mkTransport :: IO Transport } tcpConfig :: TransportConfig tcpConfig = TransportConfig - { transportName = "TCP" - , mkTransport = do + { transportName = "TCP", + mkTransport = do Right t <- TCP.createTransport (TCP.defaultTCPAddr "127.0.0.1" "0") TCP.defaultTCPParameters pure t } @@ -48,8 +49,8 @@ tcpConfig = quicConfig :: TransportConfig quicConfig = TransportConfig - { transportName = "QUIC" - , mkTransport = + { transportName = "QUIC", + mkTransport = QUIC.credentialLoadX509 -- Generate a self-signed x509v3 certificate using this nifty tool: -- https://certificatetools.com/ @@ -59,34 +60,41 @@ quicConfig = Left errmsg -> throwIO $ userError errmsg Right credentials -> QUIC.createTransport - ( QUIC.QUICTransportConfig - { hostName = "127.0.0.1" - , serviceName = "0" - , credentials = credentials :| [] - , -- credentials are self-signed - validateCredentials = False + ( (QUIC.defaultQUICTransportConfig "127.0.0.1" (credentials :| [])) + { QUIC.serviceName = "0", + QUIC.validateCredentials = False, + -- For benchmarks with lots of streams and tiny messages, we can easily + -- overflow the receive buffer + QUIC.socketOptions = [(N.RecvBuffer, 4 * 1024 * 1024)] } ) } data BenchParams = BenchParams - { messageSize :: !Int - , messageCount :: !Int - , connectionCount :: !Int + { messageSize :: !Int, + messageCount :: !Int, + connectionCount :: !Int } smallMessages, mediumMessages, largeMessages :: BenchParams -smallMessages = BenchParams{messageSize = 64, messageCount = 10_000, connectionCount = 1} -mediumMessages = BenchParams{messageSize = 1024, messageCount = 1_000, connectionCount = 1} -largeMessages = BenchParams{messageSize = 4096, messageCount = 100, connectionCount = 1} +smallMessages = BenchParams {messageSize = 64, messageCount = 10_000, connectionCount = 1} +mediumMessages = BenchParams {messageSize = 1024, messageCount = 1_000, connectionCount = 1} +largeMessages = BenchParams {messageSize = 4096, messageCount = 100, connectionCount = 1} multiConn :: Int -> BenchParams -> BenchParams -multiConn n p = p{connectionCount = n} +multiConn n p = p {connectionCount = n} throughputBench :: TransportConfig -> BenchParams -> IO () -throughputBench TransportConfig{mkTransport} BenchParams{messageSize, messageCount, connectionCount} = do +throughputBench cfg params = + timeout 30_000_000 (throughputBench' cfg params) + >>= maybe (throwIO $ userError "benchmark stalled: timed out waiting for messages") pure + +throughputBench' :: TransportConfig -> BenchParams -> IO () +throughputBench' TransportConfig {mkTransport} BenchParams {messageSize, messageCount, connectionCount} = do transport <- mkTransport - flip finally (closeTransport transport) $ do + -- Closing is bounded as well: it can block if a connection was lost, and that + -- would hide the failure we are trying to report. + flip finally (void $ timeout 5_000_000 (closeTransport transport)) $ do Right senderEP <- newEndPoint transport Right receiverEP <- newEndPoint transport @@ -94,61 +102,73 @@ throughputBench TransportConfig{mkTransport} BenchParams{messageSize, messageCou totalMessages = messageCount * connectionCount receiverReady <- newEmptyMVar - receiverDone <- newEmptyMVar - - void $ forkIO $ do - connsEstablished <- newIORef (0 :: Int) - let waitForConnections = do - event <- receive receiverEP - case event of - ConnectionOpened{} -> do - n <- atomicModifyIORef' connsEstablished (\x -> (x + 1, x + 1)) - when (n < connectionCount) waitForConnections - _ -> waitForConnections - waitForConnections - putMVar receiverReady () - - msgsReceived <- newIORef (0 :: Int) - let recvLoop = do - event <- receive receiverEP - case event of - Received _ _ -> do - n <- atomicModifyIORef' msgsReceived (\x -> (x + 1, x + 1)) - when (n < totalMessages) recvLoop - _ -> recvLoop - recvLoop - putMVar receiverDone () - - let receiverAddr = address receiverEP - connections <- - replicateM - connectionCount - (connect senderEP receiverAddr ReliableOrdered defaultConnectHints >>= either throwIO pure) - - takeMVar receiverReady - - forConcurrently_ connections $ \conn -> - forM_ [0 .. messageCount] $ \_ -> send conn [payload] - - takeMVar receiverDone - -benchTransport :: TransportConfig -> TestTree -benchTransport cfg@TransportConfig{transportName} = + + let receiver = do + connsEstablished <- newIORef (0 :: Int) + let waitForConnections = do + event <- receive receiverEP + case event of + ConnectionOpened {} -> do + n <- atomicModifyIORef' connsEstablished (\x -> (x + 1, x + 1)) + when (n < connectionCount) waitForConnections + ErrorEvent err -> throwIO err + _ -> waitForConnections + waitForConnections + putMVar receiverReady () + + msgsReceived <- newIORef (0 :: Int) + let recvLoop = do + event <- receive receiverEP + case event of + Received _ _ -> do + n <- atomicModifyIORef' msgsReceived (\x -> (x + 1, x + 1)) + when (n < totalMessages) recvLoop + ErrorEvent err -> throwIO err + _ -> recvLoop + recvLoop + + let watchSender = forever $ do + event <- receive senderEP + case event of + ErrorEvent err -> throwIO err + _ -> pure () + + withAsync receiver $ \receiverAsync -> withAsync watchSender $ \senderAsync -> do + link receiverAsync + link senderAsync + + let receiverAddr = address receiverEP + connections <- + replicateM + connectionCount + (connect senderEP receiverAddr ReliableOrdered defaultConnectHints >>= either throwIO pure) + + takeMVar receiverReady + + forConcurrently_ connections $ \conn -> + forM_ [0 .. messageCount] $ \_ -> send conn [payload] >>= either throwIO pure + + wait receiverAsync + +benchTransport :: TransportConfig -> Benchmark +benchTransport cfg@TransportConfig {transportName} = bgroup transportName [ bgroup "throughput" [ bgroup "single-connection" - [ bench "small-msg" $ nfIO $ throughputBench cfg smallMessages - , bench "default-msg" $ nfIO $ throughputBench cfg mediumMessages - , bench "large-msg" $ nfIO $ throughputBench cfg largeMessages - ] - , bgroup + [ bench "small-msg" $ nfIO $ throughputBench cfg smallMessages, + bench "default-msg" $ nfIO $ throughputBench cfg mediumMessages, + bench "large-msg" $ nfIO $ throughputBench cfg largeMessages + ], + bgroup "multi-connection" - [ bench "2-conn" $ nfIO $ throughputBench cfg smallMessages{connectionCount = 2, messageCount = 10_000} - , bench "5-conn" $ nfIO $ throughputBench cfg smallMessages{connectionCount = 5, messageCount = 10_000} - , bench "10-conn" $ nfIO $ throughputBench cfg smallMessages{connectionCount = 10, messageCount = 5_000} + [ bench "2-conn" $ nfIO $ throughputBench cfg smallMessages {connectionCount = 2, messageCount = 10_000}, + bench "5-conn" $ nfIO $ throughputBench cfg smallMessages {connectionCount = 5, messageCount = 10_000}, + bench "10-conn" $ nfIO $ throughputBench cfg smallMessages {connectionCount = 10, messageCount = 5_000}, + bench "50-conn" $ nfIO $ throughputBench cfg smallMessages {connectionCount = 50, messageCount = 100}, + bench "100-conn" $ nfIO $ throughputBench cfg smallMessages {connectionCount = 100, messageCount = 50} ] ] ] @@ -156,6 +176,9 @@ benchTransport cfg@TransportConfig{transportName} = main :: IO () main = defaultMain - [ benchTransport tcpConfig - , benchTransport quicConfig + -- QUIC is a userspace networking protocol, + -- so CPU time isn't the appropriate comparison + -- to make with TCP + [ localOption WallTime (benchTransport tcpConfig), + localOption WallTime (benchTransport quicConfig) ] diff --git a/packages/network-transport-quic/network-transport-quic.cabal b/packages/network-transport-quic/network-transport-quic.cabal index 883dec8d..350affc0 100644 --- a/packages/network-transport-quic/network-transport-quic.cabal +++ b/packages/network-transport-quic/network-transport-quic.cabal @@ -1,6 +1,6 @@ cabal-version: 3.0 Name: network-transport-quic -Version: 0.1.2 +Version: 0.2.0 build-Type: Simple License: BSD-3-Clause License-file: LICENSE @@ -59,10 +59,9 @@ library , microlens-platform ^>=0.4 , network >= 3.1 && < 3.3 , network-transport >= 0.5 && < 0.6 - -- Prior to version 0.2.20, `quic` had issues with handling - -- pending data in the stream buffer. This meant that vectored - -- message sends did not work correctly at the transport layer - , quic >=0.2.20 && <0.4 + -- Version 0.3.15 added graceful server shutdown which + -- changes the way network-transport-quic works + , quic >=0.3.15 && <0.4 , stm >=2.4 && <2.6 , tls >= 2.1 && < 2.5 , tls-session-manager >= 0.0.5 && <0.2 @@ -97,6 +96,7 @@ test-suite network-transport-quic-tests , network-transport , network-transport-quic , network-transport-tests + , quic , tasty ^>=1.5 , tasty-flaky ^>= 0.1.3 , tasty-hedgehog @@ -108,13 +108,15 @@ benchmark network-transport-quic-bench hs-source-dirs: bench main-is: Bench.hs default-language: Haskell2010 - ghc-options: -rtsopts -with-rtsopts=-N + -- -T makes the allocations of each benchmark appear in its results + ghc-options: -rtsopts "-with-rtsopts=-N -T" build-depends: async , base >=4.14 && <5 , bytestring , filepath + , network , network-transport , network-transport-tcp , network-transport-quic - , tasty ^>=1.5 + , tasty ^>=1.5 , tasty-bench >=0.4 diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal.hs index 76bd5cbd..3482fe13 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal.hs @@ -18,6 +18,9 @@ module Network.Transport.QUIC.Internal decodeMessage, MessageReceived (..), encodeMessage, + + -- * Handshake + handshake, ) where @@ -29,7 +32,7 @@ import Control.Concurrent.STM.TQueue readTQueue, writeTQueue, ) -import Control.Exception (Exception (displayException), IOException, bracket, throwIO, try) +import Control.Exception (Exception (displayException, fromException), SomeAsyncException, SomeException, bracket, catch, finally, throwIO, try) import Control.Monad (unless, when) import Data.Bifunctor (Bifunctor (first)) import Data.Binary qualified as Binary (decodeOrFail) @@ -64,7 +67,8 @@ import Network.Transport.QUIC.Internal.Messaging createConnectionId, decodeMessage, encodeMessage, - receiveMessage, + handshake, + messageReceiver, recvWord32, sendAck, sendCloseConnection, @@ -84,6 +88,7 @@ import Network.Transport.QUIC.Internal.QUICTransport TransportState (..), ValidRemoteEndPointState (..), closeLocalEndpoint, + closeLocalEndpointDeferred, closeRemoteEndPoint, createConnectionTo, createRemoteEndPoint, @@ -105,7 +110,7 @@ import Network.Transport.QUIC.Internal.QUICTransport transportState, (^.), ) -import Network.Transport.QUIC.Internal.Server (forkServer) +import Network.Transport.QUIC.Internal.Server (forkServer, stopServer) -- | Create a new Transport based on the QUIC protocol. -- @@ -118,7 +123,7 @@ createTransport initialConfig = do quicTransport <- newQUICTransport initialConfig let resolvedConfig = quicTransport ^. transportConfig - serverThread <- + server <- forkServer (quicTransport ^. transportInputSocket) (credentials resolvedConfig) @@ -129,12 +134,12 @@ createTransport initialConfig = do pure $ Transport { newEndPoint = newTQueueIO >>= newEndpoint quicTransport, - closeTransport = - foldOpenEndPoints quicTransport (closeLocalEndpoint quicTransport) - >> killThread serverThread -- TODO: use a synchronization mechanism to close the thread gracefully - >> modifyMVar_ - (quicTransport ^. transportState) - (\_ -> pure TransportStateClosed) + closeTransport = do + shutdownPeers <- foldOpenEndPoints quicTransport (closeLocalEndpointDeferred quicTransport) + stopServer server `finally` sequence_ shutdownPeers + modifyMVar_ + (quicTransport ^. transportState) + (\_ -> pure TransportStateClosed) } -- | Handle a new incoming connection. @@ -177,6 +182,7 @@ handleNewStream quicTransport stream = do (remoteEndPoint, _) <- either throwIO pure =<< createRemoteEndPoint ourEndPoint remoteAddress Incoming doneMVar <- newEmptyMVar + drained <- newEmptyMVar let serverConnId = remoteServerConnId remoteEndPoint -- One logical connection per stream; clientConnId is always 0. @@ -186,7 +192,8 @@ handleNewStream quicTransport stream = do RemoteEndPointValid $ ValidRemoteEndPointState { _remoteStream = stream, - _remoteStreamIsClosed = doneMVar + _remoteStreamIsClosed = doneMVar, + _remoteStreamDrained = drained } modifyMVar_ (remoteEndPoint ^. remoteEndPointState) @@ -220,10 +227,15 @@ handleNewStream quicTransport stream = do handleIncomingMessages ourEndPoint remoteEndPoint + `finally` tryPutMVar doneMVar () - takeMVar doneMVar - QUIC.shutdownStream stream - killThread tid + -- Once 'ConnectionClosed' (or the like) was enqueued, finishing our + -- end of the stream tells the other end that it was. + ( takeMVar doneMVar + >> (QUIC.shutdownStream stream `catch` \(_ :: SomeException) -> pure ()) + >> killThread tid + ) + `finally` tryPutMVar drained () -- | Infinite loop that listens for messages from the remote endpoint and processes them. -- @@ -237,39 +249,44 @@ handleIncomingMessages ourEndPoint remoteEndPoint = remoteAddress = remoteEndPoint ^. remoteEndPointAddress remoteState = remoteEndPoint ^. remoteEndPointState - acquire :: IO (Either IOError QUIC.Stream) + acquire :: IO (Either String QUIC.Stream) acquire = withMVar remoteState $ \case - RemoteEndPointInit -> pure . Left $ userError "handleIncomingMessages (init)" - RemoteEndPointClosed -> pure . Left $ userError "handleIncomingMessages (closed)" + RemoteEndPointInit -> pure . Left $ "handleIncomingMessages (init)" + RemoteEndPointClosed -> pure . Left $ "handleIncomingMessages (closed)" RemoteEndPointValid validState -> pure . Right $ validState ^. remoteStream - release :: Either IOError QUIC.Stream -> IO () - release (Left err) = closeRemoteEndPoint Incoming remoteEndPoint >> prematureExit err + release :: Either String QUIC.Stream -> IO () + release (Left reason) = connectionLost reason release (Right _) = closeRemoteEndPoint Incoming remoteEndPoint -- One logical connection per stream; clientConnId is always 0. connectionId = createConnectionId serverConnId 0 - go = either prematureExit loop + go = either (const $ pure ()) run + + run stream = + (messageReceiver stream >>= loop) + `catch` \(exc :: SomeException) -> case fromException exc of + Just (_ :: SomeAsyncException) -> throwIO exc + Nothing -> connectionLost (displayException exc) - loop stream = - receiveMessage stream + loop nextMessage = + nextMessage >>= \case - Left errmsg -> do - -- Throwing will trigger 'prematureExit' - throwIO $ userError $ "(handleIncomingMessages) Failed with: " <> errmsg - Right (Message bytes) -> handleMessage bytes >> loop stream - Right StreamClosed -> throwIO $ userError "(handleIncomingMessages) Stream closed" + Left errmsg -> connectionLost $ "(handleIncomingMessages) Failed with: " <> errmsg + Right (Message bytes) -> handleMessage bytes >> loop nextMessage + Right StreamClosed -> connectionLost "(handleIncomingMessages) Stream closed" Right CloseConnection -> do - atomically (writeTQueue ourQueue (ConnectionClosed connectionId)) mAct <- modifyMVar (remoteEndPoint ^. remoteEndPointState) $ \case RemoteEndPointInit -> pure (RemoteEndPointClosed, Nothing) RemoteEndPointClosed -> pure (RemoteEndPointClosed, Nothing) - RemoteEndPointValid (ValidRemoteEndPointState _ isClosed) -> do + RemoteEndPointValid (ValidRemoteEndPointState _ isClosed _) -> do pure (RemoteEndPointClosed, Just $ putMVar isClosed ()) case mAct of Nothing -> pure () - Just cleanup -> cleanup + Just cleanup -> do + atomically (writeTQueue ourQueue (ConnectionClosed connectionId)) + cleanup Right CloseEndPoint -> do -- handleIncomingMessages only runs on incoming remote endpoints, so if -- the state was still Valid there is exactly one logical connection to @@ -278,28 +295,30 @@ handleIncomingMessages ourEndPoint remoteEndPoint = RemoteEndPointValid _ -> pure (RemoteEndPointClosed, True) other -> pure (other, False) when wasValid $ - atomically $ writeTQueue ourQueue (ConnectionClosed connectionId) + atomically $ + writeTQueue ourQueue (ConnectionClosed connectionId) handleMessage :: [ByteString] -> IO () handleMessage payload = atomically (writeTQueue ourQueue (Received connectionId payload)) - prematureExit :: IOException -> IO () - prematureExit exc = do - modifyMVar_ remoteState $ \case - RemoteEndPointValid {} -> pure RemoteEndPointClosed - RemoteEndPointInit -> pure RemoteEndPointClosed - RemoteEndPointClosed -> pure RemoteEndPointClosed - atomically - ( writeTQueue - ourQueue - ( ErrorEvent - ( TransportError - (EventConnectionLost remoteAddress) - (displayException exc) - ) - ) - ) + connectionLost :: String -> IO () + connectionLost reason = do + wasValid <- modifyMVar remoteState $ \case + RemoteEndPointValid {} -> pure (RemoteEndPointClosed, True) + RemoteEndPointInit -> pure (RemoteEndPointClosed, False) + RemoteEndPointClosed -> pure (RemoteEndPointClosed, False) + when wasValid $ + atomically + ( writeTQueue + ourQueue + ( ErrorEvent + ( TransportError + (EventConnectionLost remoteAddress) + reason + ) + ) + ) newEndpoint :: QUICTransport -> @@ -379,7 +398,7 @@ newConnection ourEndPoint creds validateCreds remoteAddress _reliability _connec True -> pure . Left $ TransportError SendFailed "Remote endpoint closed" closeConn remoteEndPoint connAlive = do mCleanup <- modifyMVar (remoteEndPoint ^. remoteEndPointState) $ \case - RemoteEndPointValid vst@(ValidRemoteEndPointState stream isClosed) -> do + RemoteEndPointValid vst@(ValidRemoteEndPointState stream isClosed _) -> do readIORef connAlive >>= \case False -> pure (RemoteEndPointValid vst, Nothing) True -> do diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Client.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Client.hs index 7ccc1d4e..c9e1484c 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Client.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Client.hs @@ -1,111 +1,135 @@ {-# LANGUAGE LambdaCase #-} +{-# LANGUAGE NumericUnderscores #-} {-# LANGUAGE ScopedTypeVariables #-} -{-# LANGUAGE TupleSections #-} -{-# LANGUAGE TypeApplications #-} -module Network.Transport.QUIC.Internal.Client ( - streamToEndpoint, -) +module Network.Transport.QUIC.Internal.Client + ( PeerConnection (..), + connectToPeer, + openStream, + superviseStream, + closeTimeout, + ) where -import Control.Concurrent (forkIOWithUnmask, newEmptyMVar) -import Control.Concurrent.Async (withAsync) -import Control.Concurrent.MVar (MVar, putMVar, takeMVar, tryPutMVar) -import Control.Exception (SomeException, bracket, catch, finally, mask, mask_, throwIO) +import Control.Concurrent (forkIO) +import Control.Concurrent.Async (wait, withAsync) +import Control.Concurrent.MVar (MVar, newEmptyMVar, putMVar, takeMVar, tryPutMVar) +import Control.Exception (SomeAsyncException, SomeException, catch, displayException, finally, fromException, mask_, throwIO, try) +import Control.Monad (void) import Data.List.NonEmpty (NonEmpty) import Network.QUIC qualified as QUIC import Network.QUIC.Client qualified as QUIC.Client -import Network.Transport (ConnectErrorCode (ConnectNotFound), EndPointAddress, TransportError (..)) +import Network.Transport (ConnectErrorCode (ConnectFailed, ConnectNotFound), EndPointAddress, TransportError (..)) import Network.Transport.QUIC.Internal.Configuration (Credential, mkClientConfig) -import Network.Transport.QUIC.Internal.Messaging (MessageReceived (..), handshake, receiveMessage) +import Network.Transport.QUIC.Internal.Messaging (MessageReceived (..), closeTimeout, handshake, receiveMessage) import Network.Transport.QUIC.Internal.QUICAddr (QUICAddr (QUICAddr), decodeQUICAddr) +import System.Timeout (timeout) -streamToEndpoint :: +data PeerConnection = PeerConnection + { peerQUICConnection :: !QUIC.Connection, + peerShutdown :: !(MVar ()) + } + +-- | Like 'try', but asynchronous exceptions (cancellation, timeouts) propagate. +tryAny :: IO a -> IO (Either SomeException a) +tryAny act = + try act >>= \case + Left exc | Just (_ :: SomeAsyncException) <- fromException exc -> throwIO exc + other -> pure other + +-- | Establish a QUIC connection to the host of the given endpoint. +connectToPeer :: NonEmpty Credential -> -- | Validate credentials Bool -> - -- | Our address - EndPointAddress -> -- | Their address EndPointAddress -> - -- | Called when the QUIC connection or stream ends without us having initiated the - -- close. Must be idempotent (the caller typically gates on remote endpoint state so - -- that repeated invocations are safe) — this handler is invoked from multiple sites - -- (peer-initiated close signal, QUIC.Client.run exception, thread finally) to cover - -- every termination path. + -- | Called exactly once when the QUIC connection is gone, whatever the reason + -- (including having failed to establish it). Must not block. IO () -> - IO - ( Either - (TransportError ConnectErrorCode) - ( MVar () - , -- \^ put '()' to close the stream - QUIC.Stream - ) - ) -streamToEndpoint creds validateCreds ourAddress theirAddress onConnLoss = + IO (Either (TransportError ConnectErrorCode) PeerConnection) +connectToPeer creds validateCreds theirAddress onLost = case decodeQUICAddr theirAddress of Left errmsg -> pure $ Left (TransportError ConnectNotFound errmsg) Right (QUICAddr hostname servicename _) -> do clientConfig <- mkClientConfig hostname servicename creds validateCreds - streamMVar <- newEmptyMVar - doneMVar <- newEmptyMVar + connMVar <- newEmptyMVar + shutdown <- newEmptyMVar - let runClient :: QUIC.Connection -> IO () - runClient conn = mask $ \restore -> do - QUIC.waitEstablished conn - restore $ - bracket (QUIC.stream conn) QUIC.closeStream $ \stream -> do - handshake (ourAddress, theirAddress) stream - >>= either - (\_ -> putMVar streamMVar (Left $ TransportError ConnectNotFound "handshake failed")) - (\_ -> putMVar streamMVar (Right stream)) + let failed :: String -> IO () + failed msg = void $ tryPutMVar connMVar (Left $ TransportError ConnectNotFound msg) - withAsync (listenForClose stream doneMVar) $ \_ -> - takeMVar doneMVar + _ <- + forkIO $ + ( ( QUIC.Client.run clientConfig $ \conn -> do + QUIC.waitEstablished conn + putMVar connMVar (Right $ PeerConnection conn shutdown) + takeMVar shutdown + ) + `catch` (\(exc :: SomeException) -> failed (displayException exc)) + ) + `finally` (failed "connection closed" >> onLost) - _ <- mask_ $ - forkIOWithUnmask $ - \unmask -> - catch - ( unmask $ - QUIC.Client.run - clientConfig - ( \conn -> - catch - (runClient conn) - (throwIO @SomeException) - ) - ) - (\(_ :: SomeException) -> pure ()) - `finally` onConnLoss + takeMVar connMVar - streamOrError <- takeMVar streamMVar +openStream :: + PeerConnection -> + -- | Our address + EndPointAddress -> + -- | Their address + EndPointAddress -> + IO (Either (TransportError ConnectErrorCode) QUIC.Stream) +openStream peer ourAddress theirAddress = + tryAny (QUIC.stream (peerQUICConnection peer)) >>= \case + Left exc -> pure $ Left (TransportError ConnectFailed (displayException exc)) + Right stream -> + tryAny (handshake (ourAddress, theirAddress) stream) >>= \case + Right (Right ()) -> pure (Right stream) + Right (Left ()) -> abandon stream >> pure (Left (TransportError ConnectNotFound "handshake failed")) + Left exc -> abandon stream >> pure (Left (TransportError ConnectFailed (displayException exc))) + where + abandon = void . tryAny . QUIC.closeStream + +superviseStream :: + QUIC.Stream -> + -- | Put '()' to request that the stream be closed + MVar () -> + -- | Filled when the stream is closed + MVar () -> + -- | Called when the stream ends without us having asked for it. + IO () -> + -- | Called when the stream is finished with + IO () -> + IO () +superviseStream stream closeRequested drained onConnLoss onFinished = + void . forkIO $ + withAsync listenForClose (\listener -> takeMVar closeRequested >> drain listener) + `finally` (void (timeout closeTimeout (tryAny (QUIC.closeStream stream))) >> tryPutMVar drained () >> onFinished) + where + drain listener = + void . timeout closeTimeout . tryAny $ do + QUIC.shutdownStream stream + wait listener - pure $ (doneMVar,) <$> streamOrError - where - listenForClose :: QUIC.Stream -> MVar () -> IO () - listenForClose stream doneMVar = - receiveMessage stream - >>= \case - -- Any message from the peer on this stream means we're done listening. - -- Peer-initiated closes (StreamClosed/CloseEndPoint) additionally call - -- onConnLoss; the idempotent gate in the handler dedupes with the finally - -- that also fires on QUIC.Client.run exit. - -- - -- Mask signalling+onConnLoss as an atomic pair: tryPutMVar unblocks - -- runClient's takeMVar, which causes withAsync to cancel this thread. - -- Without mask, the async ThreadKilled can fire partway through - -- onConnLoss, dropping the ErrorEvent. The finally in the parent thread - -- is a backup but cannot recover if surfaceConnectionLost already - -- transitioned the remote state to Closed. - Right StreamClosed -> mask_ $ do - _ <- tryPutMVar doneMVar () - onConnLoss - Right CloseConnection -> - -- Peer closed the logical connection cleanly; no ErrorEvent. - () <$ tryPutMVar doneMVar () - Right CloseEndPoint -> mask_ $ do - _ <- tryPutMVar doneMVar () - onConnLoss - other -> throwIO . userError $ "Unexpected incoming message to client: " <> show other + listenForClose :: IO () + listenForClose = + ( receiveMessage stream + >>= \case + -- Peer-initiated closes (StreamClosed/CloseEndPoint) additionally call + -- onConnLoss; its idempotent gate dedupes with other termination paths. + -- + -- Mask signalling+onConnLoss as an atomic pair: tryPutMVar unblocks + -- the thread which cancels us. Without mask, the cancellation could fire + -- partway through onConnLoss, dropping the ErrorEvent. + Right StreamClosed -> lost + Right CloseConnection -> + -- Peer closed the logical connection cleanly; no ErrorEvent. + void $ tryPutMVar closeRequested () + Right CloseEndPoint -> lost + other -> throwIO . userError $ "Unexpected incoming message to client: " <> show other + ) + `catch` \(exc :: SomeException) -> case fromException exc of + Just (_ :: SomeAsyncException) -> throwIO exc + Nothing -> lost -- e.g. the QUIC connection failed + lost = mask_ $ tryPutMVar closeRequested () >> onConnLoss diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Configuration.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Configuration.hs index 16e8fd26..f21de142 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Configuration.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Configuration.hs @@ -1,46 +1,48 @@ - - -module Network.Transport.QUIC.Internal.Configuration ( - mkClientConfig, +module Network.Transport.QUIC.Internal.Configuration + ( mkClientConfig, mkServerConfig, -- * Re-export to generate credentials Credential, TLS.credentialLoadX509, -) where + ) +where import Data.List.NonEmpty (NonEmpty) import Data.List.NonEmpty qualified as NonEmpty -import Network.QUIC.Client (ClientConfig(ccValidate), ccPortName, ccServerName, defaultClientConfig) -import Network.QUIC.Internal (ServerConfig, ccCredentials) +import Network.QUIC.Client (ClientConfig (ccValidate), ccPortName, ccServerName, defaultClientConfig) +import Network.QUIC.Internal (Parameters (initialMaxStreamsBidi), ServerConfig (scParameters), ccCredentials, defaultParameters) import Network.QUIC.Server (ServerConfig (scCredentials, scSessionManager), defaultServerConfig) import Network.Socket (HostName, ServiceName) import Network.TLS (Credential, Credentials (Credentials)) import Network.Transport.QUIC.Internal.TLS qualified as TLS mkClientConfig :: - HostName -> - ServiceName -> - NonEmpty Credential -> - Bool -> -- ^ Validate credentials - IO ClientConfig + HostName -> + ServiceName -> + NonEmpty Credential -> + -- | Validate credentials + Bool -> + IO ClientConfig mkClientConfig host port creds validate = do - pure $ - defaultClientConfig - { ccServerName = host - , ccPortName = port - , ccValidate = validate - , ccCredentials = Credentials (NonEmpty.toList creds) - } + pure $ + defaultClientConfig + { ccServerName = host, + ccPortName = port, + ccValidate = validate, + ccCredentials = Credentials (NonEmpty.toList creds) + } mkServerConfig :: - NonEmpty Credential -> - IO ServerConfig + NonEmpty Credential -> + IO ServerConfig mkServerConfig creds = do - tlsSessionManager <- TLS.sessionManager + tlsSessionManager <- TLS.sessionManager - pure $ - defaultServerConfig - { scSessionManager = tlsSessionManager - , scCredentials = Credentials (NonEmpty.toList creds) - } + pure $ + defaultServerConfig + { scSessionManager = tlsSessionManager, + scCredentials = Credentials (NonEmpty.toList creds), + -- We support lots of streams per connection for dense network topologies + scParameters = defaultParameters {initialMaxStreamsBidi = 65536} + } diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Messaging.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Messaging.hs index 019154b0..53c9a5ae 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Messaging.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Messaging.hs @@ -15,6 +15,7 @@ module Network.Transport.QUIC.Internal.Messaging createConnectionId, sendMessage, receiveMessage, + messageReceiver, MessageReceived (..), -- * Specialized messages @@ -24,6 +25,7 @@ module Network.Transport.QUIC.Internal.Messaging recvWord32, sendCloseConnection, sendCloseEndPoint, + closeTimeout, -- * Handshake protocol handshake, @@ -34,7 +36,7 @@ module Network.Transport.QUIC.Internal.Messaging ) where -import Control.Exception (SomeException, catch, displayException, mask, throwIO, try) +import Control.Exception (SomeAsyncException, SomeException, catch, displayException, fromException, mask, throwIO, try) import Control.Monad (replicateM) import Data.Binary (Binary) import Data.Binary qualified as Binary @@ -42,6 +44,7 @@ import Data.Bits (shiftL, (.|.)) import Data.ByteString (ByteString) import Data.ByteString qualified as BS import Data.Functor ((<&>)) +import Data.IORef (IORef, newIORef, readIORef, writeIORef) import Data.Word (Word32, Word8) import GHC.Exception (Exception) import Network.QUIC (Stream) @@ -49,6 +52,7 @@ import Network.QUIC qualified as QUIC import Network.Transport (ConnectionId, EndPointAddress) import Network.Transport.Internal (decodeWord32, encodeWord32) import Network.Transport.QUIC.Internal.QUICAddr (QUICAddr (QUICAddr), decodeQUICAddr) +import System.Timeout (timeout) -- | Send a message on the stream. -- @@ -65,21 +69,29 @@ sendMessage stream messages = (encodeMessage messages) ) --- | Receive a message, including its local destination endpoint ID +-- | Receive a single message. -- --- This function is thread-safe; while the data is being received, asynchronous --- exceptions are masked, to be rethrown after the data is sent. +-- To receive several messages from a stream, use 'messageReceiver'. receiveMessage :: Stream -> IO (Either String MessageReceived) -receiveMessage stream = mask $ \restore -> - restore - ( decodeMessage - -- Note that 'recvStream' may return less bytes than requested. - -- Therefore, we must wrap it in 'getAllBytes'. - (getAllBytes (QUIC.recvStream stream)) - ) - `catch` (\(ex :: QUIC.QUICException) -> throwIO ex) +receiveMessage stream = messageReceiver stream >>= id + +-- | Create an action which receives the next message from a stream, every time it +-- is run. Only one such receiver should exist per stream. +messageReceiver :: + Stream -> + IO (IO (Either String MessageReceived)) +messageReceiver stream = do + -- The whole purpose of 'messageReceiver' is to amortize + -- reading with the following buffer + buffer <- newIORef BS.empty + pure $ + decodeMessage + -- Note that 'recvStream' may return less bytes than requested. + -- Therefore, we must wrap it in 'getAllBytes'. + (getAllBytes buffer (QUIC.recvStream stream)) + `catch` (\(ex :: QUIC.QUICException) -> throwIO ex) -- | Encode a message. -- @@ -107,8 +119,12 @@ decodeMessage get = >>= maybe (pure $ Right StreamClosed) ( \controlByte -> - go controlByte `catch` (\(ex :: SomeException) -> pure $ Left (displayException ex)) - ) . flip BS.indexMaybe 0 + go controlByte `catch` \(ex :: SomeException) -> + case fromException ex of + Just (_ :: SomeAsyncException) -> throwIO ex + Nothing -> pure $ Left (displayException ex) + ) + . flip BS.indexMaybe 0 where go ctrl | ctrl == closeEndPointControlByte = pure $ Right CloseEndPoint @@ -127,18 +143,31 @@ decodeMessage get = -- fetcher that repeatedly returns empty after a peer FIN would cause this to -- spin forever. getAllBytes :: + -- | Bytes fetched, but not yet consumed + IORef ByteString -> -- | Function to fetch at most 'n' bytes (Int -> IO ByteString) -> -- | Function to fetch exactly 'n' bytes (or fewer on EOF) (Int -> IO ByteString) -getAllBytes get n = go n mempty +getAllBytes buffer get n = do + buffered <- readIORef buffer + go [buffered] (BS.length buffered) where - go 0 !acc = pure $ BS.concat acc - go m !acc = - get m >>= \bytes -> - if BS.null bytes - then pure $ BS.concat acc - else go (m - BS.length bytes) (acc <> [bytes]) + go !acc !have + | have >= n = do + let (wanted, rest) = BS.splitAt n (BS.concat (reverse acc)) + writeIORef buffer rest + pure wanted + | otherwise = + get (max (n - have) fetchSize) >>= \bytes -> + if BS.null bytes + then do + writeIORef buffer BS.empty + pure $ BS.concat (reverse acc) + else go (bytes : acc) (have + BS.length bytes) + + fetchSize :: Int + fetchSize = 16384 data MessageReceived = Message {-# UNPACK #-} ![ByteString] @@ -189,8 +218,7 @@ recvWord32 :: recvWord32 stream = mask $ \restore -> restore - ( QUIC.recvStream stream 4 <&> Right . decodeWord32 - ) + (QUIC.recvStream stream 4 <&> Right . decodeWord32) `catch` (\(ex :: SomeException) -> pure $ Left (displayException ex)) -- | We perform some special actions based on a message's control byte. @@ -212,24 +240,29 @@ closeEndPointControlByte = 127 closeConnectionControlByte :: ControlByte closeConnectionControlByte = 255 +-- | How long to wait for the remote end to take a message which closes a connection, +-- or to acknowledge that a stream was closed. +closeTimeout :: Int +closeTimeout = 1_000_000 + +-- | Send a control message which says that we are done with a stream. +-- +-- Closing must never wait on the remote end: if it stopped reading, or is gone +-- without us having noticed, the stream's flow control window may never reopen and +-- sending would block forever. We give up after 'closeTimeout' instead; whoever is +-- on the other side will find out when the QUIC connection ends. +sendClosing :: ControlByte -> Stream -> IO (Either QUIC.QUICException ()) +sendClosing controlByte stream = + try (timeout closeTimeout (QUIC.sendStream stream (BS.singleton controlByte))) + <&> fmap (const ()) + -- | Send a message to close the connection. sendCloseConnection :: Stream -> IO (Either QUIC.QUICException ()) -sendCloseConnection stream = - try - ( QUIC.sendStream - stream - (BS.singleton closeConnectionControlByte) - ) +sendCloseConnection = sendClosing closeConnectionControlByte --- | Send a message to close the connection. +-- | Send a message to close the endpoint. sendCloseEndPoint :: Stream -> IO (Either QUIC.QUICException ()) -sendCloseEndPoint stream = - try - ( QUIC.sendStream - stream - ( BS.singleton closeEndPointControlByte - ) - ) +sendCloseEndPoint = sendClosing closeEndPointControlByte -- | Handshake protocol that a client, connecting to a remote endpoint, -- has to perform: diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICAddr.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICAddr.hs index 299e8671..242e749b 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICAddr.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICAddr.hs @@ -1,12 +1,13 @@ {-# LANGUAGE DerivingStrategies #-} {-# LANGUAGE GeneralizedNewtypeDeriving #-} -module Network.Transport.QUIC.Internal.QUICAddr ( - EndPointId (..), +module Network.Transport.QUIC.Internal.QUICAddr + ( EndPointId (..), QUICAddr (..), encodeQUICAddr, decodeQUICAddr, -) where + ) +where import Data.Binary (Binary) import Data.ByteString.Char8 qualified as BS8 @@ -15,48 +16,46 @@ import Data.Word (Word32) import Network.Socket (HostName, ServiceName) import Network.Transport (EndPointAddress (EndPointAddress)) -{- | Represents the unique ID of an endpoint within a transport. - -This is used by endpoints to identify remote endpoints, even though -the remote endpoints are all backed by the same QUIC address. --} +-- | Represents the unique ID of an endpoint within a transport. +-- +-- This is used by endpoints to identify remote endpoints, even though +-- the remote endpoints are all backed by the same QUIC address. newtype EndPointId = EndPointId Word32 - deriving newtype (Eq, Show, Ord, Read, Bounded, Enum, Real, Integral, Num, Binary) + deriving newtype (Eq, Show, Ord, Read, Bounded, Enum, Real, Integral, Num, Binary) -- A QUICAddr represents the unique address an `endpoint` has, which involves -- pointing to the transport (HostName, ServiceName) and then specific -- endpoint spawned by that transport (EndpointId) data QUICAddr = QUICAddr - { quicBindHost :: !HostName - , quicBindPort :: !ServiceName - , quicEndpointId :: !EndPointId - } - deriving (Eq, Ord, Show) + { quicBindHost :: !HostName, + quicBindPort :: !ServiceName, + quicEndpointId :: !EndPointId + } + deriving (Eq, Ord, Show) -- | Encode a 'QUICAddr' to 'EndPointAddress' encodeQUICAddr :: QUICAddr -> EndPointAddress encodeQUICAddr (QUICAddr host port ix) = - EndPointAddress - (BS8.pack $ host <> ":" <> port <> ":" <> show ix) + EndPointAddress + (BS8.pack $ host <> ":" <> port <> ":" <> show ix) -- | Decode end point address decodeQUICAddr :: - EndPointAddress -> - Either String QUICAddr + EndPointAddress -> + Either String QUICAddr decodeQUICAddr (EndPointAddress bs) = - case splitMaxFromEnd (== ':') 2 $ BSC.unpack bs of - [host, port, endPointIdStr] -> - case reads endPointIdStr of - [(endPointId, "")] -> Right $ QUICAddr host port endPointId - _ -> Left $ "Undecodeable 'EndPointAddress': " <> show bs - _ -> - Left $ "Undecodeable 'EndPointAddress': " <> show bs - -{- | @spltiMaxFromEnd p n xs@ splits list @xs@ at elements matching @p@, -returning at most @p@ segments -- counting from the /end/ + case splitMaxFromEnd (== ':') 2 $ BSC.unpack bs of + [host, port, endPointIdStr] -> + case reads endPointIdStr of + [(endPointId, "")] -> Right $ QUICAddr host port endPointId + _ -> Left $ "Undecodeable 'EndPointAddress': " <> show bs + _ -> + Left $ "Undecodeable 'EndPointAddress': " <> show bs -> splitMaxFromEnd (== ':') 2 "ab:cd:ef:gh" == ["ab:cd", "ef", "gh"] --} +-- | @spltiMaxFromEnd p n xs@ splits list @xs@ at elements matching @p@, +-- returning at most @p@ segments -- counting from the /end/ +-- +-- > splitMaxFromEnd (== ':') 2 "ab:cd:ef:gh" == ["ab:cd", "ef", "gh"] splitMaxFromEnd :: (a -> Bool) -> Int -> [a] -> [[a]] splitMaxFromEnd p = \n -> go [[]] n . reverse where @@ -64,7 +63,7 @@ splitMaxFromEnd p = \n -> go [[]] n . reverse go accs _ [] = accs go ([] : accs) 0 xs = reverse xs : accs go (acc : accs) n (x : xs) = - if p x - then go ([] : acc : accs) (n - 1) xs - else go ((x : acc) : accs) n xs + if p x + then go ([] : acc : accs) (n - 1) xs + else go ((x : acc) : accs) n xs go _ _ _ = error "Bug in splitMaxFromEnd" diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICTransport.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICTransport.hs index 1a3376bc..4faea835 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICTransport.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/QUICTransport.hs @@ -4,6 +4,7 @@ {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TemplateHaskell #-} +{-# LANGUAGE TupleSections #-} {-# LANGUAGE TypeApplications #-} module Network.Transport.QUIC.Internal.QUICTransport @@ -34,14 +35,19 @@ module Network.Transport.QUIC.Internal.QUICTransport nextSelfConnOutId, newLocalEndPoint, closeLocalEndpoint, + closeLocalEndpointDeferred, -- * LocalEndPointState LocalEndPointState (..), ValidLocalEndPointState, incomingConnections, outgoingConnections, + outgoingPeers, nextConnectionCounter, + -- ** OutgoingPeer + OutgoingPeer, + -- ** ConnectionCounter ConnectionCounter, @@ -60,6 +66,7 @@ module Network.Transport.QUIC.Internal.QUICTransport ValidRemoteEndPointState (..), remoteStream, remoteStreamIsClosed, + remoteStreamDrained, Direction (..), -- * Re-exports @@ -67,17 +74,22 @@ module Network.Transport.QUIC.Internal.QUICTransport ) where -import Control.Concurrent.Async (forConcurrently_) -import Control.Concurrent.MVar (MVar, modifyMVar, modifyMVar_, newMVar, readMVar, tryPutMVar) +import Control.Concurrent (forkIO) +import Control.Concurrent.Async (forConcurrently) +import Control.Concurrent.MVar (MVar, modifyMVar, modifyMVar_, newEmptyMVar, newMVar, readMVar, tryPutMVar, tryReadMVar) import Control.Concurrent.STM.TQueue (TQueue, writeTQueue) -import Control.Exception (bracketOnError) -import Control.Monad (forM_) +import Control.Exception (bracketOnError, onException) +import Control.Monad (filterM, forM_, unless, void, when) import Control.Monad.STM (atomically) +import Data.Foldable (for_) import Data.Function ((&)) +import Data.Functor ((<&>)) +import Data.IORef (IORef, atomicModifyIORef', newIORef) import Data.List.NonEmpty (NonEmpty) import Data.List.NonEmpty qualified as NE import Data.Map.Strict (Map) import Data.Map.Strict qualified as Map +import Data.Maybe (catMaybes) import Data.Word (Word32) import Lens.Micro.Platform (makeLenses, (%~), (+~), (^.)) import Network.QUIC (Stream) @@ -85,7 +97,13 @@ import Network.Socket (HostName, ServiceName, Socket) import Network.Socket qualified as N import Network.TLS (Credential) import Network.Transport (ConnectErrorCode (ConnectFailed), EndPointAddress, Event (EndPointClosed, ErrorEvent), EventErrorCode (EventConnectionLost), NewEndPointErrorCode (NewEndPointFailed), TransportError (TransportError)) -import Network.Transport.QUIC.Internal.Client (streamToEndpoint) +import Network.Transport.QUIC.Internal.Client + ( PeerConnection (..), + closeTimeout, + connectToPeer, + openStream, + superviseStream, + ) import Network.Transport.QUIC.Internal.Messaging ( ClientConnId, ServerConnId, @@ -94,6 +112,7 @@ import Network.Transport.QUIC.Internal.Messaging sendCloseEndPoint, ) import Network.Transport.QUIC.Internal.QUICAddr (EndPointId, QUICAddr (..), encodeQUICAddr) +import System.Timeout (timeout) {- The QUIC transport has three levels of statefullness: @@ -123,7 +142,17 @@ data QUICTransportConfig = QUICTransportConfig -- | Note that if your credentials is self-signed, you will have -- to turn off 'validateCredentials'. This should only be set to 'False' -- in tests, or in a private network. - validateCredentials :: Bool + validateCredentials :: Bool, + -- | A list of socket options to apply to the socker underlying a connection. + -- Socket options are applied in the order that they are specified. + -- + -- Note that socket addressed are always re-used ('N.ReuseAddr'), regardless of socket options. + -- Other potentially relevant socket options include 'N.RecvBuffer' and 'N.SendBuffer'. + -- + -- Unsupported socket options are ignored. + -- + -- @since 0.2.0 + socketOptions :: [(N.SocketOption, Int)] } deriving (Eq, Show) @@ -133,7 +162,8 @@ defaultQUICTransportConfig host creds = { hostName = host, serviceName = "443", credentials = creds, - validateCredentials = True + validateCredentials = True, + socketOptions = [] } data QUICTransport = QUICTransport @@ -163,13 +193,15 @@ newQUICTransport config = do ) N.close $ \socket -> do - N.setSocketOption socket N.ReuseAddr 1 + for_ ((N.ReuseAddr, 1) : socketOptions config) $ \(opt, val) -> + N.whenSupported opt $ N.setSocketOption socket opt val + N.withFdSocket socket N.setCloseOnExecIfNeeded N.bind socket (N.addrAddress addr) port <- N.socketPort socket QUICTransport - config{serviceName=show port} + config {serviceName = show port} socket <$> newMVar (TransportStateValid $ ValidTransportState mempty 1) @@ -197,6 +229,7 @@ data LocalEndPointState data ValidLocalEndPointState = ValidLocalEndPointState { _incomingConnections :: Map (EndPointAddress, ConnectionCounter) RemoteEndPoint, _outgoingConnections :: Map (EndPointAddress, ConnectionCounter) RemoteEndPoint, + _outgoingPeers :: Map EndPointAddress OutgoingPeer, _nextSelfConnOutId :: !ClientConnId, -- | We identify connections by remote endpoint address, AND ConnectionCounter, -- to support multiple connections between the same two endpoint addresses @@ -218,18 +251,28 @@ instance Show RemoteEndPoint where show (RemoteEndPoint address _ _) = " show address <> ">" data RemoteEndPointState - = -- | In the short window between a connection - -- being initiated and the handshake completing + = -- | In the short window between a connection being initiated and the handshake completing RemoteEndPointInit | RemoteEndPointValid ValidRemoteEndPointState | RemoteEndPointClosed data ValidRemoteEndPointState = ValidRemoteEndPointState { _remoteStream :: Stream, - _remoteStreamIsClosed :: MVar () + _remoteStreamIsClosed :: MVar (), + _remoteStreamDrained :: MVar () + } + +data OutgoingPeer = OutgoingPeer + { _peerConnection :: !(MVar (Either (TransportError ConnectErrorCode) PeerConnection)), + _peerLostReported :: !(IORef Bool), + _peerStreams :: !(MVar (Maybe (Map EndPointId (RemoteEndPoint, MVar ())))) } +instance Show OutgoingPeer where + show _ = "" + makeLenses ''QUICTransport +makeLenses ''OutgoingPeer makeLenses ''TransportState makeLenses ''ValidTransportState makeLenses ''LocalEndPoint @@ -238,6 +281,26 @@ makeLenses ''ValidLocalEndPointState makeLenses ''RemoteEndPoint makeLenses ''ValidRemoteEndPointState +dropPeer :: LocalEndPoint -> EndPointAddress -> OutgoingPeer -> IO () +dropPeer localEndPoint remoteAddress peer = + modifyMVar_ (localEndPoint ^. localEndPointState) $ \case + LocalEndPointStateClosed -> pure LocalEndPointStateClosed + LocalEndPointStateValid st -> + pure . LocalEndPointStateValid $ + st & outgoingPeers %~ Map.update (\current -> if sameAs current then Nothing else Just current) remoteAddress + where + sameAs current = (current ^. peerConnection) == (peer ^. peerConnection) + +registerStream :: OutgoingPeer -> RemoteEndPoint -> MVar () -> IO Bool +registerStream peer remoteEndPoint drained = + modifyMVar (peer ^. peerStreams) $ \case + Nothing -> pure (Nothing, False) + Just current -> pure (Just (Map.insert (remoteEndPoint ^. remoteEndPointId) (remoteEndPoint, drained) current), True) + +unregisterStream :: OutgoingPeer -> RemoteEndPoint -> IO () +unregisterStream peer remoteEndPoint = + modifyMVar_ (peer ^. peerStreams) (pure . fmap (Map.delete (remoteEndPoint ^. remoteEndPointId))) + -- | Fold over all open local endpoitns of a transport foldOpenEndPoints :: QUICTransport -> (LocalEndPoint -> IO a) -> IO [a] foldOpenEndPoints quicTransport f = @@ -259,6 +322,7 @@ newLocalEndPoint quicTransport newLocalQueue = do ValidLocalEndPointState { _incomingConnections = mempty, _outgoingConnections = mempty, + _outgoingPeers = mempty, _nextConnInId = firstNonReservedServerConnId, _nextSelfConnOutId = 0, _nextConnectionCounter = 0 @@ -291,7 +355,19 @@ closeLocalEndpoint :: QUICTransport -> LocalEndPoint -> IO () -closeLocalEndpoint quicTransport localEndPoint = do +closeLocalEndpoint quicTransport localEndPoint = closeLocalEndpointDeferred quicTransport localEndPoint >>= id + +-- | Close a local endpoint, but return an action which will close the outgoing connections. +-- +-- This function is really only useful when shutting down the whole transport, +-- where we have to do some cleanup before fully closing all endpoints. +-- +-- You should prefer to use 'closeLocalEndpoint' in most cases. +closeLocalEndpointDeferred :: + QUICTransport -> + LocalEndPoint -> + IO (IO ()) +closeLocalEndpointDeferred quicTransport localEndPoint = do modifyMVar_ (quicTransport ^. transportState) $ \case TransportStateClosed -> pure TransportStateClosed TransportStateValid vst -> @@ -304,16 +380,17 @@ closeLocalEndpoint quicTransport localEndPoint = do LocalEndPointStateClosed -> pure (LocalEndPointStateClosed, Nothing) LocalEndPointStateValid st -> pure (LocalEndPointStateClosed, Just st) - -- Close outgoing remote endpoints before incoming. The peer's handleIncomingMessages - -- reader writes ConnectionClosed in response to our outgoing close; its listenForClose - -- writes ErrorEvent in response to our incoming close. Processing outgoing first gives - -- the peer's event queue the expected ConnectionClosed-before-ErrorEvent ordering. + -- Close outgoing remote endpoints before incoming forM_ mPreviousState $ \vst -> do - forConcurrently_ (vst ^. outgoingConnections) tryCloseRemoteStream - forConcurrently_ (vst ^. incomingConnections) tryCloseRemoteStream + outgoingDrained <- catMaybes <$> forConcurrently (Map.elems $ vst ^. outgoingConnections) tryCloseRemoteStream + _ <- timeout closeTimeout (mapM_ readMVar outgoingDrained) + void $ forConcurrently (Map.elems $ vst ^. incomingConnections) tryCloseRemoteStream atomically $ writeTQueue (localEndPoint ^. localQueue) EndPointClosed + -- Everything we had to say on these QUIC connections has been said. + pure $ forM_ mPreviousState $ \vst -> forM_ (vst ^. outgoingPeers) shutdownPeer where - tryCloseRemoteStream :: RemoteEndPoint -> IO () + -- Returns the MVar which is filled once the stream is closed, if we had to close it. + tryCloseRemoteStream :: RemoteEndPoint -> IO (Maybe (MVar ())) tryCloseRemoteStream remoteEndPoint = do mCleanup <- modifyMVar (remoteEndPoint ^. remoteEndPointState) $ \case RemoteEndPointInit -> pure (RemoteEndPointClosed, Nothing) @@ -324,12 +401,10 @@ closeLocalEndpoint quicTransport localEndPoint = do Just $ do _ <- sendCloseEndPoint (vst ^. remoteStream) _ <- tryPutMVar (vst ^. remoteStreamIsClosed) () - pure () + pure (vst ^. remoteStreamDrained) ) - case mCleanup of - Nothing -> pure () - Just cleanup -> cleanup + sequence mCleanup -- | Attempt to close a remote endpoint. If the remote endpoint is in -- any non-valid state (e.g. already closed), then nothing happens. @@ -341,7 +416,7 @@ closeRemoteEndPoint direction remoteEndPoint = do mAct <- modifyMVar (remoteEndPoint ^. remoteEndPointState) $ \case RemoteEndPointInit -> pure (RemoteEndPointClosed, Nothing) RemoteEndPointClosed -> pure (RemoteEndPointClosed, Nothing) - RemoteEndPointValid (ValidRemoteEndPointState stream isClosed) -> + RemoteEndPointValid (ValidRemoteEndPointState stream isClosed _) -> let cleanup = do _ <- case direction of Outgoing -> sendCloseConnection stream @@ -402,52 +477,135 @@ createConnectionTo creds validateCreds localEndPoint remoteAddress = do createRemoteEndPoint localEndPoint remoteAddress Outgoing >>= \case Left err -> pure $ Left err Right (remoteEndPoint, _) -> do - -- TODO: each call to @connect@ currently opens a dedicated QUIC connection - -- and carries a single logical connection on its stream. Preferred - -- architecture: one QUIC connection per (local endpoint, peer endpoint) - -- pair, with each logical connection carried on its own stream. Streams - -- already give us independent flow control and avoid head-of-line blocking. - streamToEndpoint - creds - validateCreds - (localEndPoint ^. localAddress) - remoteAddress - (surfaceConnectionLost remoteEndPoint) - >>= \case - Left exc -> pure $ Left exc - Right (closeStream, stream) -> do - let validState = - RemoteEndPointValid $ - ValidRemoteEndPointState - { _remoteStream = stream, - _remoteStreamIsClosed = closeStream - } - modifyMVar_ - (remoteEndPoint ^. remoteEndPointState) - (\_ -> pure validState) - pure $ Right remoteEndPoint + let abandon :: TransportError ConnectErrorCode -> IO (Either (TransportError ConnectErrorCode) a) + abandon err = do + modifyMVar_ (remoteEndPoint ^. remoteEndPointState) (\_ -> pure RemoteEndPointClosed) + pure $ Left err + + acquirePeer creds validateCreds localEndPoint remoteAddress >>= \case + Left err -> abandon err + Right (peer, peerConn) -> do + awaitPendingCloses peer + openStream peerConn (localEndPoint ^. localAddress) remoteAddress >>= \case + Left err -> abandon err + Right stream -> do + closeRequested <- newEmptyMVar + drained <- newEmptyMVar + -- The remote endpoint must be Valid before anything can observe the + -- stream ending, or a loss would be missed. + modifyMVar_ + (remoteEndPoint ^. remoteEndPointState) + (\_ -> pure . RemoteEndPointValid $ ValidRemoteEndPointState stream closeRequested drained) + + registerStream peer remoteEndPoint drained >>= \case + False -> do + -- The peer was lost while we were connecting + _ <- tryPutMVar closeRequested () + abandon (TransportError ConnectFailed "Connection lost") + True -> do + superviseStream + stream + closeRequested + drained + (surfaceConnectionLost localEndPoint remoteAddress peer remoteEndPoint) + (unregisterStream peer remoteEndPoint) + pure $ Right remoteEndPoint where - -- Idempotent: surfaces EventConnectionLost exactly once, only if the remote - -- endpoint was still Valid when invoked. Called from multiple termination - -- sites (peer-initiated close, QUIC exception, forked-thread finally) so that - -- no close path can leave us silent — the state-transition gate dedupes them. - surfaceConnectionLost remoteEndPoint = do - mAct <- modifyMVar (remoteEndPoint ^. remoteEndPointState) $ \case - RemoteEndPointInit -> pure (RemoteEndPointClosed, Nothing) - RemoteEndPointClosed -> pure (RemoteEndPointClosed, Nothing) - RemoteEndPointValid (ValidRemoteEndPointState stream isClosed) -> - let cleanup = do - _ <- sendCloseConnection stream - _ <- tryPutMVar isClosed () - onConnectionLost - in pure (RemoteEndPointClosed, Just cleanup) - case mAct of - Nothing -> pure () - Just act -> act - onConnectionLost = - atomically - . writeTQueue (localEndPoint ^. localQueue) - . ErrorEvent - $ TransportError - (EventConnectionLost remoteAddress) - "Connection reset" + awaitPendingCloses peer = do + streams <- maybe [] Map.elems <$> readMVar (peer ^. peerStreams) + closing <- flip filterM streams $ \(remoteEndPoint, _) -> + readMVar (remoteEndPoint ^. remoteEndPointState) <&> \case + RemoteEndPointValid _ -> False + _ -> True + unless (null closing) $ + () <$ timeout closeTimeout (forM_ closing (readMVar . snd)) + +acquirePeer :: + NonEmpty Credential -> + -- | Validate credentials + Bool -> + LocalEndPoint -> + EndPointAddress -> + IO (Either (TransportError ConnectErrorCode) (OutgoingPeer, PeerConnection)) +acquirePeer creds validateCreds localEndPoint remoteAddress = do + candidate <- OutgoingPeer <$> newEmptyMVar <*> newIORef False <*> newMVar (Just mempty) + + claim <- modifyMVar (localEndPoint ^. localEndPointState) $ \case + LocalEndPointStateClosed -> + pure (LocalEndPointStateClosed, Left $ TransportError ConnectFailed "endpoint is closed") + LocalEndPointStateValid st -> case Map.lookup remoteAddress (st ^. outgoingPeers) of + Just peer -> pure (LocalEndPointStateValid st, Right (peer, False)) + Nothing -> + pure + ( LocalEndPointStateValid (st & outgoingPeers %~ Map.insert remoteAddress candidate), + Right (candidate, True) + ) + + case claim of + Left err -> pure $ Left err + Right (peer, weMustConnect) -> do + when weMustConnect $ do + result <- + connectToPeer creds validateCreds remoteAddress (onPeerLost peer) + `onException` do + _ <- tryPutMVar (peer ^. peerConnection) (Left $ TransportError ConnectFailed "interrupted") + dropPeer localEndPoint remoteAddress peer + _ <- tryPutMVar (peer ^. peerConnection) result + + either (const $ dropPeer localEndPoint remoteAddress peer) (const $ pure ()) result + + stillOpen <- + readMVar (localEndPoint ^. localEndPointState) <&> \case + LocalEndPointStateValid _ -> True + LocalEndPointStateClosed -> False + unless stillOpen (shutdownPeer peer) + + fmap (peer,) <$> readMVar (peer ^. peerConnection) + where + onPeerLost peer = do + dropPeer localEndPoint remoteAddress peer + streams <- modifyMVar (peer ^. peerStreams) (\current -> pure (Nothing, maybe [] (fmap fst . Map.elems) current)) + forM_ streams (surfaceConnectionLost localEndPoint remoteAddress peer) + +-- | Idempotent: surfaces EventConnectionLost exactly once per peer, only if the remote +-- endpoint was still Valid when invoked. Called from multiple termination +-- sites (peer-initiated close, QUIC exception, loss of the QUIC connection) so that +-- no close path can leave us silent — the state-transition gate dedupes them. +surfaceConnectionLost :: LocalEndPoint -> EndPointAddress -> OutgoingPeer -> RemoteEndPoint -> IO () +surfaceConnectionLost localEndPoint remoteAddress peer remoteEndPoint = do + mAct <- modifyMVar (remoteEndPoint ^. remoteEndPointState) $ \case + RemoteEndPointInit -> pure (RemoteEndPointClosed, Nothing) + RemoteEndPointClosed -> pure (RemoteEndPointClosed, Nothing) + RemoteEndPointValid (ValidRemoteEndPointState stream isClosed _) -> + let cleanup = do + _ <- sendCloseConnection stream + _ <- tryPutMVar isClosed () + reportPeerLost + in pure (RemoteEndPointClosed, Just cleanup) + sequence_ mAct + where + reportPeerLost = do + firstReport <- atomicModifyIORef' (peer ^. peerLostReported) (\reported -> (True, not reported)) + when firstReport $ do + dropPeer localEndPoint remoteAddress peer + atomically + . writeTQueue (localEndPoint ^. localQueue) + . ErrorEvent + $ TransportError + (EventConnectionLost remoteAddress) + "Connection reset" + + shutdownPeerWhenDrained + + shutdownPeerWhenDrained = + void . forkIO $ do + streams <- maybe [] Map.elems <$> readMVar (peer ^. peerStreams) + _ <- timeout closeTimeout (forM_ streams (readMVar . snd)) + shutdownPeer peer + +-- | Close the QUIC connection to a peer. Streams on it must have been dealt with beforehand. +shutdownPeer :: OutgoingPeer -> IO () +shutdownPeer peer = + tryReadMVar (peer ^. peerConnection) >>= \case + Just (Right peerConn) -> () <$ tryPutMVar (peerShutdown peerConn) () + _ -> pure () diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Server.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Server.hs index aa2716ad..99e38d6d 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Server.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/Server.hs @@ -1,37 +1,86 @@ +{-# LANGUAGE BangPatterns #-} +{-# LANGUAGE LambdaCase #-} +{-# LANGUAGE NumericUnderscores #-} +{-# LANGUAGE RecordWildCards #-} {-# LANGUAGE ScopedTypeVariables #-} -module Network.Transport.QUIC.Internal.Server (forkServer) where +module Network.Transport.QUIC.Internal.Server (forkServer, stopServer) where -import Control.Concurrent (ThreadId, forkIOWithUnmask) -import Control.Exception (SomeException, catch, finally, mask, mask_) +import Control.Concurrent (ThreadId, forkIOWithUnmask, killThread, threadDelay) +import Control.Concurrent.MVar (MVar, modifyMVar_, newEmptyMVar, newMVar, putMVar, readMVar, takeMVar, tryPutMVar, tryReadMVar) +import Control.Exception (SomeAsyncException, SomeException, catch, finally, fromException, mask, mask_, throwIO) +import Control.Monad (filterM, unless, void) +import Data.IORef (atomicModifyIORef', newIORef, readIORef) +import Data.IntMap.Strict (IntMap) +import Data.IntMap.Strict qualified as IntMap import Data.List.NonEmpty (NonEmpty) +import GHC.Conc (ThreadStatus (..), threadStatus) import Network.QUIC qualified as QUIC +import Network.QUIC.Internal (isConnectionClosed, mainThreadId) +import Network.QUIC.Server (scInstallShutdownHandler) import Network.QUIC.Server qualified as QUIC.Server import Network.Socket (Socket) import Network.Transport.QUIC.Internal.Configuration (Credential, mkServerConfig) +import Network.Transport.QUIC.Internal.Messaging (closeTimeout) +import System.Timeout (timeout) + +data ServerHandle = ServerHandle + { serverThread :: !ThreadId, + serverStop :: !(MVar (IO ())), + serverFinished :: !(MVar ()), + serverConnections :: !(MVar [QUIC.Connection]) + } + +stopServer :: ServerHandle -> IO () +stopServer ServerHandle {..} = do + tryReadMVar serverStop >>= \case + Nothing -> pure () + Just stop -> do + _ <- timeout closeTimeout (readMVar serverConnections >>= awaitWindingDown) + stop >> void (timeout (2 * closeTimeout) (readMVar serverFinished)) + killThread serverThread + where + awaitWindingDown conns = do + closing <- filterM windingDown conns + unless (null closing) $ threadDelay 1_000 >> awaitWindingDown closing + where + windingDown conn = do + closed <- isConnectionClosed conn + if closed then isRunning (mainThreadId conn) else pure False + +isRunning :: ThreadId -> IO Bool +isRunning tid = + threadStatus tid >>= \case + ThreadFinished -> pure False + ThreadDied -> pure False + ThreadBlocked _ -> pure True + ThreadRunning -> pure True forkServer :: Socket -> NonEmpty Credential -> -- | Error handler that runs whenever an exception is thrown inside - -- the thread that accepted an incoming connection + -- the thread that accepted an incoming connection, or a thread + -- that handles one of its streams (SomeException -> IO ()) -> -- | Termination handler that runs if the server thread catches an exception (SomeException -> IO ()) -> - -- | Request handler. The stream is closed after this handler returns. + -- | Request handler. Runs once per stream; a QUIC connection may carry many. + -- The stream is closed after this handler returns. (QUIC.Stream -> IO ()) -> - IO ThreadId + IO ServerHandle forkServer socket creds errorHandler terminationHandler requestHandler = do - serverConfig <- mkServerConfig creds + baseConfig <- mkServerConfig creds + stopVar <- newEmptyMVar + finished <- newEmptyMVar + accepted <- newMVar [] + let serverConfig = baseConfig {scInstallShutdownHandler = void . tryPutMVar stopVar} let acceptConnection :: QUIC.Connection -> IO () acceptConnection conn = mask $ \restore -> do QUIC.waitEstablished conn - stream <- QUIC.acceptStream conn - - catch - (restore (requestHandler stream `finally` QUIC.closeStream stream)) - errorHandler + modifyMVar_ accepted (\conns -> (conn :) <$> filterM (isRunning . mainThreadId) conns) + restore (acceptStreams conn errorHandler requestHandler) -- We have to make sure that the exception handler is -- installed /before/ any asynchronous exception occurs. So we mask_, then @@ -39,10 +88,43 @@ forkServer socket creds errorHandler terminationHandler requestHandler = do -- unmask only inside the catch. -- -- See the documentation for `forkIOWithUnmask`. - mask_ $ - forkIOWithUnmask - ( \unmask -> - catch - (unmask $ QUIC.Server.runWithSockets [socket] serverConfig (\conn -> catch (acceptConnection conn) errorHandler)) - terminationHandler - ) + tid <- + mask_ $ + forkIOWithUnmask + ( \unmask -> + ( catch + (unmask $ QUIC.Server.runWithSockets [socket] serverConfig (\conn -> catch (acceptConnection conn) errorHandler)) + terminationHandler + ) + `finally` tryPutMVar finished () + ) + pure ServerHandle {serverThread = tid, serverStop = stopVar, serverFinished = finished, serverConnections = accepted} + +-- | Accept the streams of a connection, handling each in its own thread. +acceptStreams :: + QUIC.Connection -> + (SomeException -> IO ()) -> + (QUIC.Stream -> IO ()) -> + IO () +acceptStreams conn errorHandler requestHandler = do + handlers <- newIORef (mempty :: IntMap ThreadId) + + let loop :: Int -> IO () + loop !n = do + stream <- QUIC.acceptStream conn + mask_ $ do + registered <- newEmptyMVar + tid <- forkIOWithUnmask $ \unmask -> do + takeMVar registered + ( unmask (requestHandler stream `finally` QUIC.closeStream stream) + `catch` \(exc :: SomeException) -> case fromException exc of + -- Being cancelled because the connection ended is expected + Just (_ :: SomeAsyncException) -> throwIO exc + Nothing -> errorHandler exc + ) + `finally` atomicModifyIORef' handlers (\m -> (IntMap.delete n m, ())) + atomicModifyIORef' handlers (\m -> (IntMap.insert n tid m, ())) + putMVar registered () + loop (n + 1) + + loop 0 `finally` (readIORef handlers >>= mapM_ killThread . IntMap.elems) diff --git a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/TLS.hs b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/TLS.hs index 4ab78014..0334560c 100644 --- a/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/TLS.hs +++ b/packages/network-transport-quic/src/Network/Transport/QUIC/Internal/TLS.hs @@ -1,13 +1,14 @@ -module Network.Transport.QUIC.Internal.TLS ( - -- * TLS session manager +module Network.Transport.QUIC.Internal.TLS + ( -- * TLS session manager sessionManager, -- * Loading TLS credentials credentialLoadX509, -) where + ) +where import Network.TLS (SessionManager, credentialLoadX509) import Network.TLS.SessionManager (defaultConfig, newSessionManager) sessionManager :: IO SessionManager -sessionManager = newSessionManager defaultConfig \ No newline at end of file +sessionManager = newSessionManager defaultConfig diff --git a/packages/network-transport-quic/test/Main.hs b/packages/network-transport-quic/test/Main.hs index 93770ac9..2b48a8b4 100644 --- a/packages/network-transport-quic/test/Main.hs +++ b/packages/network-transport-quic/test/Main.hs @@ -1,16 +1,16 @@ module Main (main) where import Test.Network.Transport.QUIC qualified (tests) -import Test.Network.Transport.QUIC.Internal.QUICAddr qualified (tests) import Test.Network.Transport.QUIC.Internal.Messaging qualified (tests) +import Test.Network.Transport.QUIC.Internal.QUICAddr qualified (tests) import Test.Tasty (defaultMain, testGroup) main :: IO () main = - defaultMain $ - testGroup - "network-transport-quic" - [ Test.Network.Transport.QUIC.Internal.Messaging.tests - , Test.Network.Transport.QUIC.Internal.QUICAddr.tests - , Test.Network.Transport.QUIC.tests - ] + defaultMain $ + testGroup + "network-transport-quic" + [ Test.Network.Transport.QUIC.Internal.Messaging.tests, + Test.Network.Transport.QUIC.Internal.QUICAddr.tests, + Test.Network.Transport.QUIC.tests + ] diff --git a/packages/network-transport-quic/test/Test/Network/Transport/QUIC.hs b/packages/network-transport-quic/test/Test/Network/Transport/QUIC.hs index b20ad1c8..dd985ddb 100644 --- a/packages/network-transport-quic/test/Test/Network/Transport/QUIC.hs +++ b/packages/network-transport-quic/test/Test/Network/Transport/QUIC.hs @@ -7,21 +7,26 @@ module Test.Network.Transport.QUIC (tests) where import Control.Concurrent.MVar (newEmptyMVar, putMVar, takeMVar) import Control.Exception (bracket) -import Control.Monad (replicateM_) +import Control.Monad (forM, forM_, replicateM, replicateM_) import Data.ByteString qualified as BS +import Data.ByteString.Char8 qualified as BSC +import Data.List (sort) import Data.List.NonEmpty (NonEmpty (..)) -import Network.Transport (EndPoint (..), Event (ConnectionClosed), Reliability (..), Transport (..), close, defaultConnectHints, send) +import Network.QUIC qualified as Q +import Network.QUIC.Client qualified as Q.Client +import Network.Transport (EndPoint (..), EndPointAddress (..), Event (..), EventErrorCode (..), Reliability (..), Transport (..), TransportError (..), close, defaultConnectHints, send) import Network.Transport.QUIC (QUICTransportConfig (..)) import Network.Transport.QUIC qualified as QUIC +import Network.Transport.QUIC.Internal (QUICAddr (..), decodeQUICAddr, handshake) import Network.Transport.Tests (echoServer) import Network.Transport.Tests qualified as Tests import Network.Transport.Tests.Auxiliary (forkTry) -import Network.Transport.Tests.Expect (expectConnectionOpened, expectEq, expectReceived, expectRight) +import Network.Transport.Tests.Expect (expectConnectionClosed, expectConnectionOpened, expectEq, expectReceived, expectRight) import Network.Transport.Util (spawn) import System.FilePath (()) import System.Timeout (timeout) import Test.Tasty (TestName, TestTree, testGroup) -import Test.Tasty.Flaky (flakyTest, limitRetries, constantDelay) +import Test.Tasty.Flaky (constantDelay, flakyTest, limitRetries) import Test.Tasty.HUnit (Assertion, assertFailure, testCase, (@?=)) tests :: TestTree @@ -33,18 +38,21 @@ tests = testCaseWithTimeout "connections" $ withQUICTransport $ flip Tests.testConnections 5, testCaseWithTimeout "closeOneConnection" $ withQUICTransport $ flip Tests.testCloseOneConnection 5, testCaseWithTimeout "closeOneDirection" $ withQUICTransport $ flip Tests.testCloseOneDirection 5, - flaky $ testCaseWithTimeout "closeReopen" $ withQUICTransport $ flip Tests.testCloseReopen 5, + testCaseWithTimeout "closeReopen" $ withQUICTransport $ flip Tests.testCloseReopen 5, -- This test is flaky specifically in Github Actions flaky $ testCaseWithTimeout "parallelConnects" $ withQUICTransport $ flip Tests.testParallelConnects 5, testCaseWithTimeout "selfSend" $ withQUICTransport Tests.testSelfSend, - flaky $ testCaseWithTimeout "closeTwice" $ withQUICTransport $ flip Tests.testCloseTwice 1, + testCaseWithTimeout "closeTwice" $ withQUICTransport $ flip Tests.testCloseTwice 1, testCaseWithTimeout "connectToSelf" $ withQUICTransport $ flip Tests.testConnectToSelf 5, testCaseWithTimeout "connectToSelfTwice" $ withQUICTransport $ flip Tests.testConnectToSelfTwice 5, testCaseWithTimeout "closeSelf" $ withQUICTransport (Tests.testCloseSelf . pure . Right), - flaky $ testCaseWithTimeout "closeEndPoint" $ withQUICTransport $ flip Tests.testCloseEndPoint 1, + testCaseWithTimeout "closeEndPoint" $ withQUICTransport $ flip Tests.testCloseEndPoint 1, flaky $ testCaseWithTimeout "closeTransport" $ Tests.testCloseTransport mkQUICTransport, - flaky $ testCaseWithTimeout "connectClosedEndPoint" $ withQUICTransport Tests.testConnectClosedEndPoint, - flaky testSendVeryLargeMessages + testCaseWithTimeout "connectClosedEndPoint" $ withQUICTransport Tests.testConnectClosedEndPoint, + testCase "Send very large messages" $ withQUICTransport testSendVeryLargeMessages, + testCaseWithTimeout "many concurrent connections to one endpoint" $ withQUICTransport testManyConnections, + testCaseWithTimeout "a connection is closed before the next is opened" $ withQUICTransport testCloseThenConnect, + testCaseWithTimeout "losing the remote end of an incoming connection is reported" $ withQUICTransport testIncomingConnectionLost ] flaky :: TestTree -> TestTree @@ -52,9 +60,13 @@ flaky = flakyTest (limitRetries 3 <> constantDelay 1_000) -- | Ensure that a test does not run for too long testCaseWithTimeout :: TestName -> Assertion -> TestTree -testCaseWithTimeout name assertion = +testCaseWithTimeout = testCaseWithTimeoutOf 1_000_000 + +-- | Like 'testCaseWithTimeout', with a timeout in microseconds. +testCaseWithTimeoutOf :: Int -> TestName -> Assertion -> TestTree +testCaseWithTimeoutOf microseconds name assertion = testCase name $ - timeout 1_000_000 assertion + timeout microseconds assertion >>= maybe (assertFailure "Test timed out") pure mkQUICTransport :: IO (Either String Transport) @@ -69,11 +81,11 @@ mkQUICTransport = do Right creds -> Right <$> QUIC.createTransport - ( QUICTransportConfig - { hostName = "127.0.0.1", - serviceName = "0", - credentials = creds :| [], - -- credentials are self-signed + ( ( QUIC.defaultQUICTransportConfig + "127.0.0.1" + (creds :| []) + ) + { serviceName = "0", validateCredentials = False } ) @@ -84,8 +96,8 @@ withQUICTransport = (mkQUICTransport >>= either assertFailure pure) closeTransport -testSendVeryLargeMessages :: TestTree -testSendVeryLargeMessages = testCase "Send very large messages" $ withQUICTransport $ \transport -> do +testSendVeryLargeMessages :: Transport -> IO () +testSendVeryLargeMessages transport = do server <- spawn transport echoServer result <- newEmptyMVar @@ -107,8 +119,77 @@ testSendVeryLargeMessages = testCase "Send very large messages" $ withQUICTransp _ <- send conn [message] (cid', payload) <- expectReceived =<< receive endpoint expectEq "connection id" cid cid' - expectEq "payload" [message] payload + expectEq "payload" [message] payload close conn receive endpoint >>= (@?=) (ConnectionClosed cid) + +testManyConnections :: Transport -> IO () +testManyConnections transport = do + let numConnections = 200 + + sender <- expectRight "newEndPoint (sender)" =<< newEndPoint transport + receiver <- expectRight "newEndPoint (receiver)" =<< newEndPoint transport + + connected <- forM [1 .. numConnections :: Int] $ \i -> do + result <- newEmptyMVar + _ <- forkTry $ do + conn <- expectRight "connect" =<< connect sender (address receiver) ReliableOrdered defaultConnectHints + expectRight "send" =<< send conn [BSC.pack (show i)] + putMVar result conn + pure result + conns <- mapM takeMVar connected + + events <- replicateM (2 * numConnections) (receive receiver) + + -- Every connection is opened before anything is received on it + let ordered _ [] = True + ordered opened (ConnectionOpened cid _ _ : rest) = ordered (cid : opened) rest + ordered opened (Received cid _ : rest) = cid `elem` opened && ordered opened rest + ordered opened (_ : rest) = ordered opened rest + expectEq "events are ordered" True (ordered [] events) + + expectEq "payloads" (sort [BSC.pack (show i) | i <- [1 .. numConnections]]) (sort [p | Received _ [p] <- events]) + + forM_ conns close + closed <- replicateM numConnections (receive receiver) + expectEq "all connections are closed" numConnections (length [() | ConnectionClosed _ <- closed]) + +testCloseThenConnect :: Transport -> IO () +testCloseThenConnect transport = do + sender <- expectRight "newEndPoint (sender)" =<< newEndPoint transport + receiver <- expectRight "newEndPoint (receiver)" =<< newEndPoint transport + + replicateM_ 100 $ do + a <- expectRight "connect (a)" =<< connect sender (address receiver) ReliableOrdered defaultConnectHints + close a + b <- expectRight "connect (b)" =<< connect sender (address receiver) ReliableOrdered defaultConnectHints + close b + + (cidA, _, _) <- expectConnectionOpened =<< receive receiver + closedA <- expectConnectionClosed =<< receive receiver + expectEq "a is closed first" cidA closedA + (cidB, _, _) <- expectConnectionOpened =<< receive receiver + closedB <- expectConnectionClosed =<< receive receiver + expectEq "then b is closed" cidB closedB + +testIncomingConnectionLost :: Transport -> IO () +testIncomingConnectionLost transport = do + receiver <- expectRight "newEndPoint" =<< newEndPoint transport + QUICAddr host port _ <- either assertFailure pure (decodeQUICAddr (address receiver)) + + let clientAddress = EndPointAddress "client" + clientConfig = Q.Client.defaultClientConfig {Q.Client.ccServerName = host, Q.Client.ccPortName = port, Q.Client.ccValidate = False} + + Q.Client.run clientConfig $ \conn -> do + Q.waitEstablished conn + stream <- Q.stream conn + handshake (clientAddress, address receiver) stream >>= either (const $ assertFailure "handshake failed") pure + + (_, _, from) <- expectConnectionOpened =<< receive receiver + expectEq "connection is from the client" clientAddress from + + receive receiver >>= \case + ErrorEvent (TransportError (EventConnectionLost lost) _) -> expectEq "the lost connection is the client's" clientAddress lost + other -> assertFailure $ "Expected the connection to be reported lost, but got " <> show other diff --git a/packages/network-transport-quic/test/Test/Network/Transport/QUIC/Internal/Messaging.hs b/packages/network-transport-quic/test/Test/Network/Transport/QUIC/Internal/Messaging.hs index 2465b20b..5999b2e1 100644 --- a/packages/network-transport-quic/test/Test/Network/Transport/QUIC/Internal/Messaging.hs +++ b/packages/network-transport-quic/test/Test/Network/Transport/QUIC/Internal/Messaging.hs @@ -15,33 +15,33 @@ import Test.Tasty.Hedgehog (testProperty) tests :: TestTree tests = - testGroup - "Messaging" - [testMessageEncodingAndDecoding] + testGroup + "Messaging" + [testMessageEncodingAndDecoding] testMessageEncodingAndDecoding :: TestTree testMessageEncodingAndDecoding = testProperty "Encoded messages can be decoded" $ property $ do - -- The message length is encoded and decoded as a Word32. Generate data above - -- a Word8 (255) to exercise the Word32 parsing of the number of bytes in each - -- message. - messages <- forAll (Gen.list (Range.linear 0 3) (Gen.bytes (Range.linear 1 4096))) - let encoded = mconcat $ encodeMessage messages + -- The message length is encoded and decoded as a Word32. Generate data above + -- a Word8 (255) to exercise the Word32 parsing of the number of bytes in each + -- message. + messages <- forAll (Gen.list (Range.linear 0 3) (Gen.bytes (Range.linear 1 4096))) + let encoded = mconcat $ encodeMessage messages - getBytes <- liftIO $ messageDecoder encoded + getBytes <- liftIO $ messageDecoder encoded - decoded <- liftIO $ decodeMessage getBytes - Right (Message messages) === decoded + decoded <- liftIO $ decodeMessage getBytes + Right (Message messages) === decoded messageDecoder :: ByteString -> IO (Int -> IO ByteString) messageDecoder allBytes = do - ref <- newIORef allBytes - pure - ( \nbytes -> do - atomicModifyIORef - ref - ( \remainingBytes -> - ( BS.drop nbytes remainingBytes - , BS.take nbytes remainingBytes - ) - ) - ) + ref <- newIORef allBytes + pure + ( \nbytes -> do + atomicModifyIORef + ref + ( \remainingBytes -> + ( BS.drop nbytes remainingBytes, + BS.take nbytes remainingBytes + ) + ) + )