{-# LANGUAGE PatternSignatures, CPP #-}
{-# OPTIONS_GHC -w #-}

module Control.Concurrent.Parallel3(parallel_, parallelStop) where

import GHC.Conc
import Control.Concurrent
import Control.Monad
import Control.Exception as E
import System.IO.Unsafe


-- initialise on the main thread, and keep
mainThread :: ThreadId
mainThread = unsafePerformIO $ myThreadId

-- True = kill the thread after it finishes
{-# NOINLINE queue #-}
queue :: Chan (IO Bool)
queue = seq mainThread $ unsafePerformIO $ do
    chan <- newChan
    replicateM_ (numCapabilities-1) $ addWorker
    return chan

#if __GLASGOW_HASKELL__ >= 610
type AnyException = SomeException
#else
type AnyException = Exception
#endif


addWorker :: IO ()
addWorker = do
    forkIO $ f `E.catch` \(e :: AnyException) -> do
        putStrLn $ "Exception on thread: " ++ show e
        throwTo mainThread $ ErrorCall $ "Control.Concurrent.Parallel: parallel thread died.\n" ++ show e
    return ()
    where
        f :: IO ()
        f = do
            kill <- join $ readChan queue
            unless kill f


-- If you don't call this then no one holds the queue, the queue gets
-- GC'd, the threads find themselves blocked indefinately, and you get
-- exceptions. This cleanly shuts down the threads, then the queue isn't important.
-- Only call this AFTER all parallel_ calls have completed.
parallelStop :: IO ()
parallelStop = replicateM_ (numCapabilities-1) $ writeChan queue $ return True


-- | Run the list of computations in parallel
--   Rule: No thread should get pre-empted (although not a guarantee)
--         On return all actions have been performed
parallel_ :: [IO a] -> IO ()
parallel_ xs | numCapabilities <= 1 = sequence_ xs
parallel_ [] = return ()
parallel_ [x] = x >> return ()
parallel_ (x1:xs) = do
    count <- newMVar $ length xs
    pause <- newEmptyMVar
    forM_ xs $ \x ->
        writeChan queue $ do
            x
            modifyMVar count $ \i -> do
                let i2 = i - 1
                    kill = i2 == 0
                when kill $ putMVar pause ()
                return (i2, kill)
    x1
    addWorker
    takeMVar pause
