{-# LANGUAGE NoFieldSelectors #-}

{-| Assemble a pipeline's descriptor-set-layout bindings and push-constant ranges
from the per-stage contributions of several shaders.

A binding (or push-constant range) declared by more than one stage must become a
single entry whose 'Vk.stageFlags' is the OR of the contributing stages — that is
what a pipeline layout shared between, say, a vertex and a fragment shader needs.
'mergeDescriptorSetLayoutBindings' and 'mergePushConstantRanges' do that merge on
plain Vulkan values, independent of where the per-stage bindings came from (hand
written, or reflected — see @vulkan-utils-spirv@).
-}
module Vulkan.Utils.PipelineLayout
  ( mergeDescriptorSetLayoutBindings
  , mergePushConstantRanges
  , DescriptorBindingConflict (..)
  ) where

import Control.Monad (foldM)
import Data.Bits ((.|.))
import Data.Foldable (foldl')
import qualified Data.Map.Strict as Map
import Data.Word (Word32)
import qualified Vulkan.Core10 as Vk
import Vulkan.Zero (zero)

{- | Two stages declared the same binding number with different descriptor types,
which cannot be reconciled into one binding.
-}
data DescriptorBindingConflict = DescriptorBindingConflict
  { DescriptorBindingConflict -> Word32
binding :: Word32
  -- ^ The binding number the stages disagree on.
  , DescriptorBindingConflict -> (DescriptorType, DescriptorType)
types :: (Vk.DescriptorType, Vk.DescriptorType)
  -- ^ The two differing descriptor types.
  }
  deriving (DescriptorBindingConflict -> DescriptorBindingConflict -> Bool
(DescriptorBindingConflict -> DescriptorBindingConflict -> Bool)
-> (DescriptorBindingConflict -> DescriptorBindingConflict -> Bool)
-> Eq DescriptorBindingConflict
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: DescriptorBindingConflict -> DescriptorBindingConflict -> Bool
== :: DescriptorBindingConflict -> DescriptorBindingConflict -> Bool
$c/= :: DescriptorBindingConflict -> DescriptorBindingConflict -> Bool
/= :: DescriptorBindingConflict -> DescriptorBindingConflict -> Bool
Eq, Int -> DescriptorBindingConflict -> ShowS
[DescriptorBindingConflict] -> ShowS
DescriptorBindingConflict -> String
(Int -> DescriptorBindingConflict -> ShowS)
-> (DescriptorBindingConflict -> String)
-> ([DescriptorBindingConflict] -> ShowS)
-> Show DescriptorBindingConflict
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> DescriptorBindingConflict -> ShowS
showsPrec :: Int -> DescriptorBindingConflict -> ShowS
$cshow :: DescriptorBindingConflict -> String
show :: DescriptorBindingConflict -> String
$cshowList :: [DescriptorBindingConflict] -> ShowS
showList :: [DescriptorBindingConflict] -> ShowS
Show)

{- | Merge the descriptor-set-layout bindings contributed by several stages for a
single descriptor set. Bindings sharing a binding number are combined: their
'Vk.stageFlags' are OR-ed and their 'Vk.descriptorCount's maxed. A
'Vk.descriptorType' disagreement is a 'Left'. The result is ascending by
binding number.

Each input binding should carry the one stage that declares it (its
'Vk.stageFlags' set to that stage); the merge turns the per-stage bindings
into one multi-stage binding per binding number.
-}
mergeDescriptorSetLayoutBindings
  :: (Foldable f)
  => f Vk.DescriptorSetLayoutBinding
  -> Either DescriptorBindingConflict [Vk.DescriptorSetLayoutBinding]
mergeDescriptorSetLayoutBindings :: forall (f :: * -> *).
Foldable f =>
f DescriptorSetLayoutBinding
-> Either DescriptorBindingConflict [DescriptorSetLayoutBinding]
mergeDescriptorSetLayoutBindings f DescriptorSetLayoutBinding
bindings =
  (Map Word32 DescriptorSetLayoutBinding
 -> [DescriptorSetLayoutBinding])
-> Either
     DescriptorBindingConflict (Map Word32 DescriptorSetLayoutBinding)
-> Either DescriptorBindingConflict [DescriptorSetLayoutBinding]
forall a b.
(a -> b)
-> Either DescriptorBindingConflict a
-> Either DescriptorBindingConflict b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (((Word32, DescriptorSetLayoutBinding)
 -> DescriptorSetLayoutBinding)
-> [(Word32, DescriptorSetLayoutBinding)]
-> [DescriptorSetLayoutBinding]
forall a b. (a -> b) -> [a] -> [b]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (Word32, DescriptorSetLayoutBinding) -> DescriptorSetLayoutBinding
forall a b. (a, b) -> b
snd ([(Word32, DescriptorSetLayoutBinding)]
 -> [DescriptorSetLayoutBinding])
-> (Map Word32 DescriptorSetLayoutBinding
    -> [(Word32, DescriptorSetLayoutBinding)])
-> Map Word32 DescriptorSetLayoutBinding
-> [DescriptorSetLayoutBinding]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Map Word32 DescriptorSetLayoutBinding
-> [(Word32, DescriptorSetLayoutBinding)]
forall k a. Map k a -> [(k, a)]
Map.toAscList) ((Map Word32 DescriptorSetLayoutBinding
 -> DescriptorSetLayoutBinding
 -> Either
      DescriptorBindingConflict (Map Word32 DescriptorSetLayoutBinding))
-> Map Word32 DescriptorSetLayoutBinding
-> f DescriptorSetLayoutBinding
-> Either
     DescriptorBindingConflict (Map Word32 DescriptorSetLayoutBinding)
forall (t :: * -> *) (m :: * -> *) b a.
(Foldable t, Monad m) =>
(b -> a -> m b) -> b -> t a -> m b
foldM Map Word32 DescriptorSetLayoutBinding
-> DescriptorSetLayoutBinding
-> Either
     DescriptorBindingConflict (Map Word32 DescriptorSetLayoutBinding)
step Map Word32 DescriptorSetLayoutBinding
forall k a. Map k a
Map.empty f DescriptorSetLayoutBinding
bindings)
  where
    step :: Map Word32 DescriptorSetLayoutBinding
-> DescriptorSetLayoutBinding
-> Either
     DescriptorBindingConflict (Map Word32 DescriptorSetLayoutBinding)
step Map Word32 DescriptorSetLayoutBinding
acc DescriptorSetLayoutBinding
b = case Word32
-> Map Word32 DescriptorSetLayoutBinding
-> Maybe DescriptorSetLayoutBinding
forall k a. Ord k => k -> Map k a -> Maybe a
Map.lookup DescriptorSetLayoutBinding
b.binding Map Word32 DescriptorSetLayoutBinding
acc of
      Maybe DescriptorSetLayoutBinding
Nothing -> Map Word32 DescriptorSetLayoutBinding
-> Either
     DescriptorBindingConflict (Map Word32 DescriptorSetLayoutBinding)
forall a b. b -> Either a b
Right (Word32
-> DescriptorSetLayoutBinding
-> Map Word32 DescriptorSetLayoutBinding
-> Map Word32 DescriptorSetLayoutBinding
forall k a. Ord k => k -> a -> Map k a -> Map k a
Map.insert DescriptorSetLayoutBinding
b.binding DescriptorSetLayoutBinding
b Map Word32 DescriptorSetLayoutBinding
acc)
      Just DescriptorSetLayoutBinding
b0 -> (\DescriptorSetLayoutBinding
b' -> Word32
-> DescriptorSetLayoutBinding
-> Map Word32 DescriptorSetLayoutBinding
-> Map Word32 DescriptorSetLayoutBinding
forall k a. Ord k => k -> a -> Map k a -> Map k a
Map.insert DescriptorSetLayoutBinding
b.binding DescriptorSetLayoutBinding
b' Map Word32 DescriptorSetLayoutBinding
acc) (DescriptorSetLayoutBinding
 -> Map Word32 DescriptorSetLayoutBinding)
-> Either DescriptorBindingConflict DescriptorSetLayoutBinding
-> Either
     DescriptorBindingConflict (Map Word32 DescriptorSetLayoutBinding)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> DescriptorSetLayoutBinding
-> DescriptorSetLayoutBinding
-> Either DescriptorBindingConflict DescriptorSetLayoutBinding
forall {r}.
(HasField "descriptorCount" r Word32,
 HasField "descriptorType" r DescriptorType,
 HasField "stageFlags" r ShaderStageFlags) =>
DescriptorSetLayoutBinding
-> r -> Either DescriptorBindingConflict DescriptorSetLayoutBinding
combine DescriptorSetLayoutBinding
b0 DescriptorSetLayoutBinding
b

    combine :: DescriptorSetLayoutBinding
-> r -> Either DescriptorBindingConflict DescriptorSetLayoutBinding
combine DescriptorSetLayoutBinding
b0 r
b1
      | DescriptorType
t0 DescriptorType -> DescriptorType -> Bool
forall a. Eq a => a -> a -> Bool
/= DescriptorType
t1 = DescriptorBindingConflict
-> Either DescriptorBindingConflict DescriptorSetLayoutBinding
forall a b. a -> Either a b
Left (Word32
-> (DescriptorType, DescriptorType) -> DescriptorBindingConflict
DescriptorBindingConflict DescriptorSetLayoutBinding
b0.binding (DescriptorType
t0, DescriptorType
t1))
      | Bool
otherwise =
          DescriptorSetLayoutBinding
-> Either DescriptorBindingConflict DescriptorSetLayoutBinding
forall a b. b -> Either a b
Right
            ( DescriptorSetLayoutBinding
b0
                { Vk.stageFlags = b0.stageFlags .|. b1.stageFlags
                , Vk.descriptorCount = max b0.descriptorCount b1.descriptorCount
                }
            )
      where
        t0 :: DescriptorType
t0 = DescriptorSetLayoutBinding
b0.descriptorType
        t1 :: DescriptorType
t1 = r
b1.descriptorType

{- | Merge the push-constant ranges contributed by several stages: ranges sharing
the same @(offset, size)@ have their 'Vk.stageFlags' OR-ed. The result is
ascending by offset.
-}
mergePushConstantRanges
  :: (Foldable f) => f Vk.PushConstantRange -> [Vk.PushConstantRange]
mergePushConstantRanges :: forall (f :: * -> *).
Foldable f =>
f PushConstantRange -> [PushConstantRange]
mergePushConstantRanges f PushConstantRange
ranges =
  [ PushConstantRange
forall a. Zero a => a
zero{Vk.stageFlags = stage, Vk.offset = off, Vk.size = sz}
  | ((Word32
off, Word32
sz), ShaderStageFlags
stage) <- Map (Word32, Word32) ShaderStageFlags
-> [((Word32, Word32), ShaderStageFlags)]
forall k a. Map k a -> [(k, a)]
Map.toAscList Map (Word32, Word32) ShaderStageFlags
byRange
  ]
  where
    byRange :: Map (Word32, Word32) ShaderStageFlags
byRange = (Map (Word32, Word32) ShaderStageFlags
 -> PushConstantRange -> Map (Word32, Word32) ShaderStageFlags)
-> Map (Word32, Word32) ShaderStageFlags
-> f PushConstantRange
-> Map (Word32, Word32) ShaderStageFlags
forall b a. (b -> a -> b) -> b -> f a -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' Map (Word32, Word32) ShaderStageFlags
-> PushConstantRange -> Map (Word32, Word32) ShaderStageFlags
forall {a} {b} {a} {r}.
(Ord a, Ord b, Bits a, HasField "size" r b, HasField "offset" r a,
 HasField "stageFlags" r a) =>
Map (a, b) a -> r -> Map (a, b) a
add Map (Word32, Word32) ShaderStageFlags
forall k a. Map k a
Map.empty f PushConstantRange
ranges
    add :: Map (a, b) a -> r -> Map (a, b) a
add Map (a, b) a
m r
r = (a -> a -> a) -> (a, b) -> a -> Map (a, b) a -> Map (a, b) a
forall k a. Ord k => (a -> a -> a) -> k -> a -> Map k a -> Map k a
Map.insertWith a -> a -> a
forall a. Bits a => a -> a -> a
(.|.) (r
r.offset, r
r.size) r
r.stageFlags Map (a, b) a
m