packages feed

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

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

module Network.Transport.QUIC.Internal.Server (forkServer, stopServer) where

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, 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. Runs once per stream; a QUIC connection may carry many.
  -- The stream is closed after this handler returns.
  (QUIC.Stream -> IO ()) ->
  IO ServerHandle
forkServer socket creds errorHandler terminationHandler requestHandler = do
  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
        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
  -- forkIOWithUnmask (the child thread inherits the masked state from the parent), then
  -- unmask only inside the catch.
  --
  -- See the documentation for `forkIOWithUnmask`.
  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)