network-transport-quic 0.1.2 → 0.2.0
raw patch · 15 files changed
+894/−462 lines, 15 filesdep ~networkdep ~quicPVP ok
version bump matches the API change (PVP)
Dependency ranges changed: network, quic
API changes (from Hackage documentation)
+ Network.Transport.QUIC: [socketOptions] :: QUICTransportConfig -> [(SocketOption, Int)]
+ Network.Transport.QUIC.Internal: [socketOptions] :: QUICTransportConfig -> [(SocketOption, Int)]
+ Network.Transport.QUIC.Internal: handshake :: (EndPointAddress, EndPointAddress) -> Stream -> IO (Either () ())
- Network.Transport.QUIC: QUICTransportConfig :: HostName -> ServiceName -> NonEmpty Credential -> Bool -> QUICTransportConfig
+ Network.Transport.QUIC: QUICTransportConfig :: HostName -> ServiceName -> NonEmpty Credential -> Bool -> [(SocketOption, Int)] -> QUICTransportConfig
- Network.Transport.QUIC.Internal: QUICTransportConfig :: HostName -> ServiceName -> NonEmpty Credential -> Bool -> QUICTransportConfig
+ Network.Transport.QUIC.Internal: QUICTransportConfig :: HostName -> ServiceName -> NonEmpty Credential -> Bool -> [(SocketOption, Int)] -> QUICTransportConfig
Files
- CHANGELOG.md +8/−0
- README.md +1/−1
- bench/Bench.hs +104/−81
- network-transport-quic.cabal +9/−7
- src/Network/Transport/QUIC/Internal.hs +67/−48
- src/Network/Transport/QUIC/Internal/Client.hs +108/−84
- src/Network/Transport/QUIC/Internal/Configuration.hs +29/−27
- src/Network/Transport/QUIC/Internal/Messaging.hs +70/−37
- src/Network/Transport/QUIC/Internal/QUICAddr.hs +32/−33
- src/Network/Transport/QUIC/Internal/QUICTransport.hs +231/−73
- src/Network/Transport/QUIC/Internal/Server.hs +101/−19
- src/Network/Transport/QUIC/Internal/TLS.hs +4/−3
- test/Main.hs +8/−8
- test/Test/Network/Transport/QUIC.hs +100/−19
- test/Test/Network/Transport/QUIC/Internal/Messaging.hs +22/−22
CHANGELOG.md view
@@ -1,3 +1,11 @@+2026-10-06 Laurent P. René de Cotret <laurent.decotret@outlook.com> 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 <laurent.decotret@outlook.com> 0.1.2
README.md view
@@ -8,7 +8,7 @@ * 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
bench/Bench.hs view
@@ -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 @@ 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 @@ 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 @@ 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 ()+ 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- _ -> recvLoop- recvLoop- putMVar receiverDone ()+ 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 receiverAddr = address receiverEP- connections <-- replicateM- connectionCount- (connect senderEP receiverAddr ReliableOrdered defaultConnectHints >>= either throwIO pure)+ let watchSender = forever $ do+ event <- receive senderEP+ case event of+ ErrorEvent err -> throwIO err+ _ -> pure () - takeMVar receiverReady+ withAsync receiver $ \receiverAsync -> withAsync watchSender $ \senderAsync -> do+ link receiverAsync+ link senderAsync - forConcurrently_ connections $ \conn ->- forM_ [0 .. messageCount] $ \_ -> send conn [payload]+ let receiverAddr = address receiverEP+ connections <-+ replicateM+ connectionCount+ (connect senderEP receiverAddr ReliableOrdered defaultConnectHints >>= either throwIO pure) - takeMVar receiverDone+ takeMVar receiverReady -benchTransport :: TransportConfig -> TestTree-benchTransport cfg@TransportConfig{transportName} =+ 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 @@ 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) ]
network-transport-quic.cabal view
@@ -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 @@ , 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 @@ , network-transport , network-transport-quic , network-transport-tests+ , quic , tasty ^>=1.5 , tasty-flaky ^>= 0.1.3 , tasty-hedgehog@@ -108,13 +108,15 @@ 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
src/Network/Transport/QUIC/Internal.hs view
@@ -18,6 +18,9 @@ decodeMessage, MessageReceived (..), encodeMessage,++ -- * Handshake+ handshake, ) where @@ -29,7 +32,7 @@ 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 @@ createConnectionId, decodeMessage, encodeMessage,- receiveMessage,+ handshake,+ messageReceiver, recvWord32, sendAck, sendCloseConnection,@@ -84,6 +88,7 @@ TransportState (..), ValidRemoteEndPointState (..), closeLocalEndpoint,+ closeLocalEndpointDeferred, closeRemoteEndPoint, createConnectionTo, createRemoteEndPoint,@@ -105,7 +110,7 @@ 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 @@ quicTransport <- newQUICTransport initialConfig let resolvedConfig = quicTransport ^. transportConfig- serverThread <-+ server <- forkServer (quicTransport ^. transportInputSocket) (credentials resolvedConfig)@@ -129,12 +134,12 @@ 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 @@ (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 @@ RemoteEndPointValid $ ValidRemoteEndPointState { _remoteStream = stream,- _remoteStreamIsClosed = doneMVar+ _remoteStreamIsClosed = doneMVar,+ _remoteStreamDrained = drained } modifyMVar_ (remoteEndPoint ^. remoteEndPointState)@@ -220,10 +227,15 @@ 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 @@ 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 - loop stream =- receiveMessage stream+ run stream =+ (messageReceiver stream >>= loop)+ `catch` \(exc :: SomeException) -> case fromException exc of+ Just (_ :: SomeAsyncException) -> throwIO exc+ Nothing -> connectionLost (displayException exc)++ 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 @@ 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 @@ 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
src/Network/Transport/QUIC/Internal/Client.hs view
@@ -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 - 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+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++ 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
src/Network/Transport/QUIC/Internal/Configuration.hs view
@@ -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}+ }
src/Network/Transport/QUIC/Internal/Messaging.hs view
@@ -15,6 +15,7 @@ createConnectionId, sendMessage, receiveMessage,+ messageReceiver, MessageReceived (..), -- * Specialized messages@@ -24,6 +25,7 @@ recvWord32, sendCloseConnection, sendCloseEndPoint,+ closeTimeout, -- * Handshake protocol handshake,@@ -34,7 +36,7 @@ ) 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.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.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,22 +69,30 @@ (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. -- -- The encoding is composed of a header, and the payloads.@@ -107,8 +119,12 @@ >>= 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,19 +143,32 @@ -- 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] | CloseConnection@@ -189,8 +218,7 @@ 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 @@ 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:
src/Network/Transport/QUIC/Internal/QUICAddr.hs view
@@ -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 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 @@ 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"
src/Network/Transport/QUIC/Internal/QUICTransport.hs view
@@ -4,6 +4,7 @@ {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TemplateHaskell #-}+{-# LANGUAGE TupleSections #-} {-# LANGUAGE TypeApplications #-} module Network.Transport.QUIC.Internal.QUICTransport@@ -34,14 +35,19 @@ nextSelfConnOutId, newLocalEndPoint, closeLocalEndpoint,+ closeLocalEndpointDeferred, -- * LocalEndPointState LocalEndPointState (..), ValidLocalEndPointState, incomingConnections, outgoingConnections,+ outgoingPeers, nextConnectionCounter, + -- ** OutgoingPeer+ OutgoingPeer,+ -- ** ConnectionCounter ConnectionCounter, @@ -60,6 +66,7 @@ ValidRemoteEndPointState (..), remoteStream, remoteStreamIsClosed,+ remoteStreamDrained, Direction (..), -- * Re-exports@@ -67,17 +74,22 @@ ) 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 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 @@ 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 @@ -- | 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 @@ { hostName = host, serviceName = "443", credentials = creds,- validateCredentials = True+ validateCredentials = True,+ socketOptions = [] } data QUICTransport = QUICTransport@@ -163,13 +193,15 @@ ) 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 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 @@ show (RemoteEndPoint address _ _) = "<RemoteEndPoint @ " <> 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 _ = "<OutgoingPeer>"+ makeLenses ''QUICTransport+makeLenses ''OutgoingPeer makeLenses ''TransportState makeLenses ''ValidTransportState makeLenses ''LocalEndPoint@@ -238,6 +281,26 @@ 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 @@ ValidLocalEndPointState { _incomingConnections = mempty, _outgoingConnections = mempty,+ _outgoingPeers = mempty, _nextConnInId = firstNonReservedServerConnId, _nextSelfConnOutId = 0, _nextConnectionCounter = 0@@ -291,7 +355,19 @@ 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 @@ 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 @@ 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 @@ 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 @@ 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 ()
src/Network/Transport/QUIC/Internal/Server.hs view
@@ -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 @@ -- 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)
src/Network/Transport/QUIC/Internal/TLS.hs view
@@ -1,10 +1,11 @@-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)
test/Main.hs view
@@ -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+ ]
test/Test/Network/Transport/QUIC.hs view
@@ -7,21 +7,26 @@ 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 @@ 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 @@ -- | 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 @@ 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 @@ (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 @@ _ <- 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
test/Test/Network/Transport/QUIC/Internal/Messaging.hs view
@@ -15,33 +15,33 @@ 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+ )+ )+ )