#pragma once #include "Array.h" #include "DX12Buffer.h" struct ID3D12Device5; struct ID3D12GraphicsCommandList; struct D3D12_INPUT_ELEMENT_DESC; struct D3D12_ROOT_PARAMETER1; struct D3D12_CONSTANT_BUFFER_VIEW_DESC; namespace Framework { enum DX12ShaderRegister { DX12_SHADER_REGISTER_B_CONST_BUFFER = 0, DX12_SHADER_REGISTER_T_SHADER_RESOURCE = 1, DX12_SHADER_REGISTER_U_UNORDERED_ACCESS = 2, // TODO: Do we need Sampler? }; struct DX12ShaderRegisterUsage { DX12ShaderRegister registerType; int registerIndex; int spaceIndex; }; class DX12ShaderSignature : public ReferenceCounter { private: ID3D12RootSignature* signature; Array registerUsages; bool changed; public: DX12ShaderSignature(); ~DX12ShaderSignature(); /** * needs to be called for each datastructure with : register(...) * * \param registerType the register type e.g. * DX12_SHADER_REGISTER_B_CONST_BUFFER for : register(b0) * \param registerIndex the register index e.g. 1 for : register(b1) */ void addRegisterUsage( DX12ShaderRegister registerType, int registerIndex); /** * needs to be called for each datastructure with : register(...) * * \param registerType the register type e.g. * DX12_SHADER_REGISTER_B_CONST_BUFFER for : register(b0, space1) * \param registerIndex the register index e.g. 1 for : register(b1, * space2) * \param spaceIndex the space index e.g. 3 for : register(b1, space3) */ void addRegisterUsage( DX12ShaderRegister registerType, int registerIndex, int spaceIndex); /** * Creates the root signature. */ void createSignature(ID3D12Device5* zDevice, PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature); ID3D12RootSignature* zSignature() const; }; class DX12ShaderFunction : public ReferenceCounter { private: Text functionName; DX12ShaderSignature* signature; D3D12_EXPORT_DESC* exportDesc; public: DX12ShaderFunction( const Text& functionName, DX12ShaderSignature* signature); ~DX12ShaderFunction(); const Text& getFunctionName() const; DX12ShaderSignature* zSignature() const; D3D12_EXPORT_DESC* zExportDesc() const; }; class DX12Shader : public ReferenceCounter { private: RCArray functions; const char* shaderBytes; int shaderBytesSize; D3D12_DXIL_LIBRARY_DESC* libraryDesc; public: DX12Shader(const char* shaderBytes, int shaderBytesSize); ~DX12Shader(); void addFunction(DX12ShaderFunction* function); int getShaderBytesSize() const; const char* getShaderBytes() const; const RCArray& getFunctions() const; D3D12_DXIL_LIBRARY_DESC* zLibraryDesc() const; }; class DX12ShaderHitGroup : public ReferenceCounter { private: Text name; DX12ShaderFunction* closestHitShaderFunction; DX12ShaderFunction* anyHitShaderFunction; DX12ShaderFunction* intersectionShaderFunction; int payloadSize; int attributeSize; D3D12_HIT_GROUP_DESC* hitGroupDesc; public: DX12ShaderHitGroup(const Text name); ~DX12ShaderHitGroup(); void setClosestHitShaderFunction( DX12ShaderFunction* closestHitShaderFunction); void setAnyHitShaderFunction(DX12ShaderFunction* anyHitShaderFunction); void setIntersectionShaderFunction( DX12ShaderFunction* intersectionShaderFunction); void setPayloadSize(int payloadSize); void setAttributeSize(int attributeSize); const Text& getName() const; DX12ShaderFunction* zClosestHitShaderFunction() const; DX12ShaderFunction* zAnyHitShaderFunction() const; DX12ShaderFunction* zIntersectionShaderFunction() const; int getPayloadSize() const; int getAttributeSize() const; D3D12_HIT_GROUP_DESC* zHitGroupDesc() const; }; class DX12Pipeline : public ReferenceCounter { private: RCArray shaders; RCArray hitGroups; ID3D12RootSignature* emptyGlobalRootSignature; ID3D12RootSignature* emptyLocalRootSignature; ID3D12StateObject* pipelineState; int maxRecursionDepth; public: DX12Pipeline(); ~DX12Pipeline(); void addShader(DX12Shader* shader); void addHitGroup(DX12ShaderHitGroup* hitGroup); void setMaxRecursionDepth(int maxRecursionDepth); void createPipelineState(ID3D12Device5* zDevice, PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature); ID3D12StateObject* zPipelineState() const; }; // namespace Framework } // namespace Framework