module Vulkan.Utils.Shader
  ( shaderStage
  , shaderModuleStage
  ) where

import Control.Monad.Trans.Resource (MonadResource, ReleaseKey, allocate)
import Data.ByteString (ByteString)
import Vulkan.CStruct.Extends (SomeStruct (..))
import qualified Vulkan.Core10 as Vk
import Vulkan.Utils.Pipeline.Specialization (Specialization, allocateSpecialization)
import Vulkan.Zero (zero)

{- | Build a 'PipelineShaderStageCreateInfo' for a single SPIR-V module with
entry point @main@. The returned 'ReleaseKey' frees the module — release it
once the pipeline is built.

The @spec@ argument supplies specialization constants (see
'Vulkan.Utils.Pipeline.Specialization'); pass @()@ for none.
-}
shaderStage
  :: (MonadResource m, Specialization spec)
  => Vk.Device
  -> Vk.ShaderStageFlagBits
  -> spec
  -> ByteString
  -> m (ReleaseKey, SomeStruct Vk.PipelineShaderStageCreateInfo)
shaderStage :: forall (m :: * -> *) spec.
(MonadResource m, Specialization spec) =>
Device
-> ShaderStageFlagBits
-> spec
-> ByteString
-> m (ReleaseKey, SomeStruct PipelineShaderStageCreateInfo)
shaderStage Device
dev ShaderStageFlagBits
stage spec
spec ByteString
code = do
  specializationInfo <- spec -> m (Maybe SpecializationInfo)
forall spec (m :: * -> *).
(Specialization spec, MonadResource m) =>
spec -> m (Maybe SpecializationInfo)
allocateSpecialization spec
spec
  shaderModuleStage dev stage specializationInfo code

{- | Lower-level companion to 'shaderStage' taking an already-built
'Vk.SpecializationInfo' (or 'Nothing'). Useful when one specialization is shared
across several stages — build it once with
'Vulkan.Utils.Pipeline.Specialization.withSpecialization' and pass it to each
stage rather than re-packing per stage.
-}
shaderModuleStage
  :: (MonadResource m)
  => Vk.Device
  -> Vk.ShaderStageFlagBits
  -> Maybe Vk.SpecializationInfo
  -> ByteString
  -> m (ReleaseKey, SomeStruct Vk.PipelineShaderStageCreateInfo)
shaderModuleStage :: forall (m :: * -> *).
MonadResource m =>
Device
-> ShaderStageFlagBits
-> Maybe SpecializationInfo
-> ByteString
-> m (ReleaseKey, SomeStruct PipelineShaderStageCreateInfo)
shaderModuleStage Device
dev ShaderStageFlagBits
stage Maybe SpecializationInfo
specializationInfo ByteString
code = do
  (key, module') <- Device
-> ShaderModuleCreateInfo '[]
-> Maybe AllocationCallbacks
-> (IO ShaderModule
    -> (ShaderModule -> IO ()) -> m (ReleaseKey, ShaderModule))
-> m (ReleaseKey, ShaderModule)
forall (a :: [*]) (io :: * -> *) r.
(Extendss ShaderModuleCreateInfo a, PokeChain a, MonadIO io) =>
Device
-> ShaderModuleCreateInfo a
-> Maybe AllocationCallbacks
-> (io ShaderModule -> (ShaderModule -> io ()) -> r)
-> r
Vk.withShaderModule Device
dev ShaderModuleCreateInfo '[]
forall a. Zero a => a
zero{Vk.code = code} Maybe AllocationCallbacks
forall a. Maybe a
Nothing IO ShaderModule
-> (ShaderModule -> IO ()) -> m (ReleaseKey, ShaderModule)
forall (m :: * -> *) a.
MonadResource m =>
IO a -> (a -> IO ()) -> m (ReleaseKey, a)
allocate
  pure
    ( key
    , SomeStruct
        zero
          { Vk.stage
          , Vk.module'
          , Vk.name = "main"
          , Vk.specializationInfo
          }
    )