module GHC.Stack.Profiler.Internal.Sampler (
  Interval (..),
  SamplerDescr (..),
  withSampler,
  startSampler,
  stopSampler,
) where

import Control.Concurrent (ThreadId, myThreadId, threadCapability, threadDelay)
import Control.Concurrent.Async (async)
import Control.Concurrent.Chan (writeChan)
import Control.Concurrent.MVar (newEmptyMVar, putMVar, takeMVar)
import Control.Exception (bracket, finally)
import Control.Monad.STM (atomically)
import qualified Control.Monad.STM as STM
import qualified Data.ByteString.Lazy as BSL
import Data.Coerce (coerce)
import Data.Foldable (for_)
import Data.Word (Word32)
import GHC.Conc (BlockReason (..), ThreadStatus (..), labelThread, threadStatus)
import GHC.Conc.Sync (fromThreadId)
import GHC.Internal.Control.Monad (forever)
import GHC.Stack.CloneStack (cloneThreadStack)
import qualified GHC.Stack.Profiler.Core as GSPC
import GHC.Stack.Profiler.Internal.Decode (
  CallStackSample (..),
  decodeToCallStack,
  serializeCallStack,
  serializeMessages,
 )
import GHC.Stack.Profiler.Internal.Manager (
  ControlMessage (..),
  Manager (..),
  Sampler (..),
  cancelSampler,
  registerSamplerThread,
  shouldProfile,
  unregisterSamplerThread,
 )

-- NOTE: Part of the public API.

-- | The sampling interval.
--
--   @since 0.5.0.0
newtype Interval
  = MkIntervalMillis {Interval -> Int
intervalMillis :: Int}
  deriving stock (Interval -> Interval -> Bool
(Interval -> Interval -> Bool)
-> (Interval -> Interval -> Bool) -> Eq Interval
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Interval -> Interval -> Bool
== :: Interval -> Interval -> Bool
$c/= :: Interval -> Interval -> Bool
/= :: Interval -> Interval -> Bool
Eq, Int -> Interval -> ShowS
[Interval] -> ShowS
Interval -> String
(Int -> Interval -> ShowS)
-> (Interval -> String) -> ([Interval] -> ShowS) -> Show Interval
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Interval -> ShowS
showsPrec :: Int -> Interval -> ShowS
$cshow :: Interval -> String
show :: Interval -> String
$cshowList :: [Interval] -> ShowS
showList :: [Interval] -> ShowS
Show)

-- | @`fromInteger` n@ constructs an interval of @n@ milliseconds.
instance Num Interval where
  (+) :: Interval -> Interval -> Interval
  + :: Interval -> Interval -> Interval
(+) = forall a b. Coercible a b => a -> b
forall a b. Coercible a b => a -> b
coerce @(Int -> Int -> Int) Int -> Int -> Int
forall a. Num a => a -> a -> a
(+)

  (-) :: Interval -> Interval -> Interval
  (-) = forall a b. Coercible a b => a -> b
forall a b. Coercible a b => a -> b
coerce @(Int -> Int -> Int) Int -> Int -> Int
forall a. Num a => a -> a -> a
(+)

  (*) :: Interval -> Interval -> Interval
  * :: Interval -> Interval -> Interval
(*) = forall a b. Coercible a b => a -> b
forall a b. Coercible a b => a -> b
coerce @(Int -> Int -> Int) Int -> Int -> Int
forall a. Num a => a -> a -> a
(*)

  abs :: Interval -> Interval
  abs :: Interval -> Interval
abs = forall a b. Coercible a b => a -> b
forall a b. Coercible a b => a -> b
coerce @(Int -> Int) Int -> Int
forall a. Num a => a -> a
abs

  signum :: Interval -> Interval
  signum :: Interval -> Interval
signum = forall a b. Coercible a b => a -> b
forall a b. Coercible a b => a -> b
coerce @(Int -> Int) Int -> Int
forall a. Num a => a -> a
signum

  fromInteger :: Integer -> Interval
  fromInteger :: Integer -> Interval
fromInteger = forall a b. Coercible a b => a -> b
forall a b. Coercible a b => a -> b
coerce @(Integer -> Int) Integer -> Int
forall a. Num a => Integer -> a
fromInteger

-- | Get the interval in microseconds.
intervalMicros :: Interval -> Int
intervalMicros :: Interval -> Int
intervalMicros = (Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
1_000) (Int -> Int) -> (Interval -> Int) -> Interval -> Int
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Interval -> Int
intervalMillis
{-# INLINE intervalMicros #-}

-- | A description used to construct a `Sampler` thread.
data SamplerDescr = MkSamplerDescr
  { SamplerDescr -> Manager
samplerManager :: Manager
  , SamplerDescr -> IO [ThreadId]
samplerThreads :: IO [ThreadId]
  , SamplerDescr -> Interval
sampleInterval :: !Interval
  }

withSampler :: SamplerDescr -> (Sampler -> IO a) -> IO a
withSampler :: forall a. SamplerDescr -> (Sampler -> IO a) -> IO a
withSampler SamplerDescr
sampler Sampler -> IO a
action =
  IO Sampler -> (Sampler -> IO ()) -> (Sampler -> IO a) -> IO a
forall a b c. IO a -> (a -> IO b) -> (a -> IO c) -> IO c
bracket
    (SamplerDescr -> IO Sampler
startSampler SamplerDescr
sampler)
    (Manager -> Sampler -> IO ()
stopSampler (SamplerDescr -> Manager
samplerManager SamplerDescr
sampler))
    Sampler -> IO a
action

-- | Run a `SamplerDescr`.
startSampler :: SamplerDescr -> IO Sampler
startSampler :: SamplerDescr -> IO Sampler
startSampler sampler :: SamplerDescr
sampler@MkSamplerDescr{Manager
samplerManager :: SamplerDescr -> Manager
samplerManager :: Manager
samplerManager, Interval
sampleInterval :: SamplerDescr -> Interval
sampleInterval :: Interval
sampleInterval} = do
  barrier <- IO (MVar ())
forall a. IO (MVar a)
newEmptyMVar
  samplerAsync <- async $ do
    () <- takeMVar barrier
    samplerThreadId <- myThreadId
    labelThread samplerThreadId $
      "Stack Sampler " <> show (fromThreadId samplerThreadId)
    forever $ do
      sampleThreads sampler
      -- TODO: Measure the delay at each step and subtract that from the next tick.
      threadDelay (intervalMicros sampleInterval)

  let
    samplerThread = MkSampler{Async ()
samplerAsync :: Async ()
samplerAsync :: Async ()
samplerAsync}

  -- Register this sampler thread to avoid sampling it
  registerSamplerThread samplerManager samplerThread
  putMVar barrier ()
  pure samplerThread

-- NOTE: `stopSampler` is part of the public API.

-- | Stop a `Sampler` thread.
--
--   @since 0.5.0.0
stopSampler :: Manager -> Sampler -> IO ()
stopSampler :: Manager -> Sampler -> IO ()
stopSampler Manager
manager Sampler
samplerThread = do
  Sampler -> IO ()
cancelSampler Sampler
samplerThread
    IO () -> IO () -> IO ()
forall a b. IO a -> IO b -> IO a
`finally` Manager -> Sampler -> IO ()
unregisterSamplerThread Manager
manager Sampler
samplerThread

-- | Take one `CallStackSample` for every thread sampled by the `Sampler`.
sampleThreads :: SamplerDescr -> IO ()
sampleThreads :: SamplerDescr -> IO ()
sampleThreads MkSamplerDescr{Manager
samplerManager :: SamplerDescr -> Manager
samplerManager :: Manager
samplerManager, IO [ThreadId]
samplerThreads :: SamplerDescr -> IO [ThreadId]
samplerThreads :: IO [ThreadId]
samplerThreads} = do
  -- Wait until the manager signals to start profiling.
  STM () -> IO ()
forall a. STM a -> IO a
atomically (Bool -> STM ()
STM.check (Bool -> STM ()) -> STM Bool -> STM ()
forall (m :: * -> *) a b. Monad m => (a -> m b) -> m a -> m b
=<< Manager -> STM Bool
shouldProfile Manager
samplerManager)
  -- List all threads that should be sampled.
  threadIds <- IO [ThreadId]
samplerThreads
  -- Sample all threads.
  for_ threadIds $ \ThreadId
threadId ->
    Manager -> ThreadId -> IO ()
sampleThread Manager
samplerManager ThreadId
threadId

-- | Take one `CallStackSample` for the given `ThreadId` and send it to the given `Manager`.
sampleThread :: Manager -> ThreadId -> IO ()
sampleThread :: Manager -> ThreadId -> IO ()
sampleThread Manager
manager ThreadId
threadId =
  ThreadId -> IO (Maybe CallStackSample)
sampleCallStackFor ThreadId
threadId
    IO (Maybe CallStackSample)
-> (Maybe CallStackSample -> IO ()) -> IO ()
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= IO ()
-> (CallStackSample -> IO ()) -> Maybe CallStackSample -> IO ()
forall b a. b -> (a -> b) -> Maybe a -> b
maybe (() -> IO ()
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()) (Manager -> CallStackSample -> IO ()
sendCallStackSample Manager
manager)

-- | Send a `CallStackSample` to the given `Manager`.
sendCallStackSample :: Manager -> CallStackSample -> IO ()
sendCallStackSample :: Manager -> CallStackSample -> IO ()
sendCallStackSample Manager
manager CallStackSample
callStackSample = do
  callStack <- CallStackSample -> IO CallStack
decodeToCallStack CallStackSample
callStackSample
  binaryMessages <-
    atomically $ do
      -- TODO: Should these two STM calls be put in a single transaction?
      messages <- serializeCallStack (symbolTableRef manager) callStack
      STM.check =<< shouldProfile manager
      pure $! serializeMessages messages
  writeChan (messageChan manager) $!
    WriteProfileSample $
      BSL.toStrict <$> binaryMessages

-- | Take a `CallStackSample` for the given `ThreadId`.
sampleCallStackFor :: ThreadId -> IO (Maybe CallStackSample)
sampleCallStackFor :: ThreadId -> IO (Maybe CallStackSample)
sampleCallStackFor ThreadId
threadId = do
  status <- ThreadId -> IO ThreadStatus
threadStatus ThreadId
threadId
  (capNo, _lockedToCap) <- threadCapability threadId
  if canTakeCallStackSample status
    then do
      cloneThreadStack threadId >>= \StackSnapshot
stackSnapshot ->
        Maybe CallStackSample -> IO (Maybe CallStackSample)
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Maybe CallStackSample -> IO (Maybe CallStackSample))
-> Maybe CallStackSample -> IO (Maybe CallStackSample)
forall a b. (a -> b) -> a -> b
$
          CallStackSample -> Maybe CallStackSample
forall a. a -> Maybe a
Just (CallStackSample -> Maybe CallStackSample)
-> CallStackSample -> Maybe CallStackSample
forall a b. (a -> b) -> a -> b
$
            CallStackSample
              { callStackSampleThreadId :: ThreadId
callStackSampleThreadId = Word64 -> ThreadId
GSPC.MkThreadId (Word64 -> ThreadId)
-> (ThreadId -> Word64) -> ThreadId -> ThreadId
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ThreadId -> Word64
fromThreadId (ThreadId -> ThreadId) -> ThreadId -> ThreadId
forall a b. (a -> b) -> a -> b
$ ThreadId
threadId
              , callStackSampleCapabilityId :: CapabilityId
callStackSampleCapabilityId = Word32 -> CapabilityId
GSPC.MkCapabilityId (Word32 -> CapabilityId) -> (Int -> Word32) -> Int -> CapabilityId
forall b c a. (b -> c) -> (a -> b) -> a -> c
. forall a b. (Integral a, Num b) => a -> b
fromIntegral @Int @Word32 (Int -> CapabilityId) -> Int -> CapabilityId
forall a b. (a -> b) -> a -> b
$ Int
capNo
              , callStackSampleStackSnapshot :: StackSnapshot
callStackSampleStackSnapshot = StackSnapshot
stackSnapshot
              }
    else pure Nothing

-- | Can a `CallStackSample` be taken for the given `ThreadId`?
canTakeCallStackSample :: ThreadStatus -> Bool
canTakeCallStackSample :: ThreadStatus -> Bool
canTakeCallStackSample = \case
  ThreadStatus
ThreadRunning -> Bool
True
  ThreadBlocked BlockReason
BlockedOnMVar -> Bool
True
  ThreadStatus
_ -> Bool
False
{-# INLINE canTakeCallStackSample #-}