packages feed

network-transport-quic-0.2.0: src/Network/Transport/QUIC/Internal/Client.hs

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

module Network.Transport.QUIC.Internal.Client
  ( PeerConnection (..),
    connectToPeer,
    openStream,
    superviseStream,
    closeTimeout,
  )
where

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 (ConnectFailed, ConnectNotFound), EndPointAddress, TransportError (..))
import Network.Transport.QUIC.Internal.Configuration (Credential, mkClientConfig)
import Network.Transport.QUIC.Internal.Messaging (MessageReceived (..), closeTimeout, handshake, receiveMessage)
import Network.Transport.QUIC.Internal.QUICAddr (QUICAddr (QUICAddr), decodeQUICAddr)
import System.Timeout (timeout)

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 ->
  -- | Their address
  EndPointAddress ->
  -- | 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) 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

      connMVar <- newEmptyMVar
      shutdown <- newEmptyMVar

      let failed :: String -> IO ()
          failed msg = void $ tryPutMVar connMVar (Left $ TransportError ConnectNotFound msg)

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

      takeMVar connMVar

openStream ::
  PeerConnection ->
  -- | Our address
  EndPointAddress ->
  -- | Their address
  EndPointAddress ->
  IO (Either (TransportError ConnectErrorCode) QUIC.Stream)
openStream peer ourAddress theirAddress =
  tryAny (QUIC.stream (peerQUICConnection peer)) >>= \case
    Left exc -> pure $ Left (TransportError ConnectFailed (displayException exc))
    Right stream ->
      tryAny (handshake (ourAddress, theirAddress) stream) >>= \case
        Right (Right ()) -> pure (Right stream)
        Right (Left ()) -> abandon stream >> pure (Left (TransportError ConnectNotFound "handshake failed"))
        Left exc -> abandon stream >> pure (Left (TransportError ConnectFailed (displayException exc)))
  where
    abandon = void . tryAny . QUIC.closeStream

superviseStream ::
  QUIC.Stream ->
  -- | Put '()' to request that the stream be closed
  MVar () ->
  -- | Filled when the stream is closed
  MVar () ->
  -- | Called when the stream ends without us having asked for it.
  IO () ->
  -- | Called when the stream is finished with
  IO () ->
  IO ()
superviseStream stream closeRequested drained onConnLoss onFinished =
  void . forkIO $
    withAsync listenForClose (\listener -> takeMVar closeRequested >> drain listener)
      `finally` (void (timeout closeTimeout (tryAny (QUIC.closeStream stream))) >> tryPutMVar drained () >> onFinished)
  where
    drain listener =
      void . timeout closeTimeout . tryAny $ do
        QUIC.shutdownStream stream
        wait listener

    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