packages feed

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 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+              )+          )+    )