packages feed

cabal-buck2-0.1.0.0: src/Distribution/Client/Buck2/Schedule.hs

-- | Running a dependency graph of jobs concurrently.
module Distribution.Client.Buck2.Schedule
  ( runDependencyGraph
  ) where

import Distribution.Client.Compat.Prelude
import Prelude ()

import qualified Data.Map as Map
import qualified Data.Set as Set

import qualified Control.Concurrent.Async as Async
import Control.Concurrent.STM
  ( atomically
  , modifyTVar'
  , newTVarIO
  , readTVar
  , retry
  , writeTVar
  )

import Distribution.Simple.Utils (ordNub)

-- | @runDependencyGraph n deps process@ runs @process k@ once for every key
-- @k@ of @deps@, on up to @n@ threads, but only after @process@ has finished
-- for each of the keys @deps@ lists for @k@. Dependencies that aren't keys of
-- @deps@ are ignored. Returns when every job has finished; if one throws, the
-- rest are cancelled and the exception is rethrown. The graph must be
-- acyclic.
--
-- Jobs are scheduled by their direct dependencies alone, so a job is
-- available to run as soon as the jobs it depends on are done.
runDependencyGraph :: Ord k => Int -> Map k [k] -> (k -> IO ()) -> IO ()
runDependencyGraph numWorkers deps process = do
  let jobs = Map.keysSet deps
      dependenciesOf = Map.map (ordNub . filter (`Set.member` jobs)) deps
      dependents = Map.fromListWith (++) [(d, [k]) | (k, ds) <- Map.toList dependenciesOf, d <- ds]
      totalJobs = Map.size deps

  remainingVar <- newTVarIO (Map.map length dependenciesOf)
  readyVar <- newTVarIO [k | (k, []) <- Map.toList dependenciesOf]
  finishedVar <- newTVarIO (0 :: Int)

  let finishJob k = do
        remaining <- readTVar remainingVar
        let (remaining', newlyReady) =
              foldl'
                ( \(rs, ready) d ->
                    let n = Map.findWithDefault 0 d rs - 1
                     in (Map.insert d n rs, if n == 0 then d : ready else ready)
                )
                (remaining, [])
                (Map.findWithDefault [] k dependents)
        writeTVar remainingVar remaining'
        modifyTVar' readyVar (newlyReady ++)
        modifyTVar' finishedVar (+ 1)

      nextJob = do
        ready <- readTVar readyVar
        case ready of
          (k : rest) -> writeTVar readyVar rest >> return (Just k)
          [] -> do
            finished <- readTVar finishedVar
            if finished == totalJobs then return Nothing else retry

      worker = do
        mk <- atomically nextJob
        for_ mk $ \k -> do
          process k
          atomically (finishJob k)
          worker

  Async.replicateConcurrently_ numWorkers worker