packages feed

network-transport-quic-0.2.0: test/Test/Network/Transport/QUIC.hs

{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE NumericUnderscores #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}

module Test.Network.Transport.QUIC (tests) where

import Control.Concurrent.MVar (newEmptyMVar, putMVar, takeMVar)
import Control.Exception (bracket)
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.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 (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 (constantDelay, flakyTest, limitRetries)
import Test.Tasty.HUnit (Assertion, assertFailure, testCase, (@?=))

tests :: TestTree
tests =
  testGroup
    "Network.Transport.QUIC"
    [ testCaseWithTimeout "ping-pong" $ withQUICTransport $ flip Tests.testPingPong 5,
      testCaseWithTimeout "endpoints" $ withQUICTransport $ flip Tests.testEndPoints 5,
      testCaseWithTimeout "connections" $ withQUICTransport $ flip Tests.testConnections 5,
      testCaseWithTimeout "closeOneConnection" $ withQUICTransport $ flip Tests.testCloseOneConnection 5,
      testCaseWithTimeout "closeOneDirection" $ withQUICTransport $ flip Tests.testCloseOneDirection 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,
      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),
      testCaseWithTimeout "closeEndPoint" $ withQUICTransport $ flip Tests.testCloseEndPoint 1,
      flaky $ testCaseWithTimeout "closeTransport" $ Tests.testCloseTransport mkQUICTransport,
      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
flaky = flakyTest (limitRetries 3 <> constantDelay 1_000)

-- | Ensure that a test does not run for too long
testCaseWithTimeout :: TestName -> Assertion -> TestTree
testCaseWithTimeout = testCaseWithTimeoutOf 1_000_000

-- | Like 'testCaseWithTimeout', with a timeout in microseconds.
testCaseWithTimeoutOf :: Int -> TestName -> Assertion -> TestTree
testCaseWithTimeoutOf microseconds name assertion =
  testCase name $
    timeout microseconds assertion
      >>= maybe (assertFailure "Test timed out") pure

mkQUICTransport :: IO (Either String Transport)
mkQUICTransport = do
  QUIC.credentialLoadX509
    -- Generate a self-signed x509v3 certificate using this nifty tool:
    -- https://certificatetools.com/
    ("test" </> "credentials" </> "cert.crt")
    ("test" </> "credentials" </> "cert.key")
    >>= \case
      Left errmsg -> pure $ Left errmsg
      Right creds ->
        Right
          <$> QUIC.createTransport
            ( ( QUIC.defaultQUICTransportConfig
                  "127.0.0.1"
                  (creds :| [])
              )
                { serviceName = "0",
                  validateCredentials = False
                }
            )

withQUICTransport :: (Transport -> IO a) -> IO a
withQUICTransport =
  bracket
    (mkQUICTransport >>= either assertFailure pure)
    closeTransport

testSendVeryLargeMessages :: Transport -> IO ()
testSendVeryLargeMessages transport = do
  server <- spawn transport echoServer
  result <- newEmptyMVar

  let numPings = 10
  let bigMessage = BS.replicate 4091 66 -- Using an odd number of bytes (4091) to test message boundaries
  _ <- forkTry $ do
    endpoint <- expectRight "newEndPoint" =<< newEndPoint transport
    ping endpoint server numPings bigMessage
    putMVar result ()

  takeMVar result
  where
    ping endpoint serverAddr numPings message = do
      conn <- expectRight "connect" =<< connect endpoint serverAddr ReliableOrdered defaultConnectHints

      (cid, _, _) <- expectConnectionOpened =<< receive endpoint

      replicateM_ numPings $ do
        _ <- send conn [message]
        (cid', payload) <- expectReceived =<< receive endpoint
        expectEq "connection id" cid cid'
        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