#pragma once
#include "backend/rhi_types.h"
#include <slang-com-helper.h>
#include <slang-com-ptr.h>
#include <slang.h>
#include <string>
#include <utility>
#include <vector>
namespace MTL {
class Allocation;
class ArgumentEncoder;
class Device;
class CommandQueue;
class CommandBuffer;
class RenderCommandEncoder;
class RenderPipelineState;
class Library;
class Function;
class Buffer;
class Texture;
class SamplerState;
class DepthStencilState;
class DepthStencilDescriptor;
class ResidencySet;
class ResidencySetDescriptor;
} // namespace MTL
namespace MTL4 {
class ArgumentTable;
class ArgumentTableDescriptor;
class CommandAllocator;
class CommandAllocatorDescriptor;
class CommandBuffer;
class CommandQueue;
class RenderCommandEncoder;
class RenderPassDescriptor;
} // namespace MTL4
namespace CA {
class MetalLayer;
class MetalDrawable;
} // namespace CA
namespace rhi {
struct Device;
struct Sampler;
struct Image;
struct CmdBuffer;
struct ShaderCompiler;
} // namespace rhi
// mt only. not rhi.
namespace mt {
struct IndexPool {
u32 cap = 0;
u32 next = 0;
std::vector<u32> free_list;
u64 alloc();
void free(u64 slot);
};
// table: 0 push, 2..17 vertex fetch. rest rides in push.
static constexpr u32 kPushTableIndex = 0;
static constexpr u32 kVertexBaseIndex = 2;
static constexpr u32 kMaxVertexBindings = 16;
static constexpr u32 kMaxBufferBindCount = kVertexBaseIndex + kMaxVertexBindings; // 18
static_assert(kMaxBufferBindCount <= 31, "MTL4 argument table allows at most 31 buffer binds");
// unused. keeps validator quiet.
static constexpr u32 kMaxTextureBindCount = 1;
static constexpr u32 kMaxSamplerBindCount = 1;
static_assert(kMaxTextureBindCount <= 128, "MTL4 argument table allows at most 128 texture binds");
static_assert(kMaxSamplerBindCount <= 16, "MTL4 argument table allows at most 16 sampler binds");
// host id lists. no gpu mirror.
static constexpr u32 kHeapTextureBudget = 8192;
static constexpr u32 kHeapSamplerBudget = 32;
struct Heap {
MTL4::ArgumentTable *table = nullptr;
MTL::ResidencySet *residency = nullptr;
// slot -> id. buffers use addresses.
std::vector<u64> texture_ids;
std::vector<u64> sampler_ids;
IndexPool sampler_space;
IndexPool texture_space;
};
std::string compile_msl(rhi::ShaderCompiler &sc, const char *path, const char *entryName);
std::string compile_msl_linked(rhi::ShaderCompiler &sc, const char *path, const char *const *entries, u32 entry_count);
void upload_image_data(rhi::Device &device, rhi::Image &image, const void *data, u32 bytes_per_row);
void create_heap(rhi::Device &device);
void destroy_heap(rhi::Device &device);
u64 write_sampler(rhi::Device &device, rhi::Sampler sampler);
u64 write_texture(rhi::Device &device, MTL::Texture *view, bool storage);
// 8b handles. slang reads them raw.
u64 heap_texture_id(rhi::Device &device, u64 slot);
u64 heap_sampler_id(rhi::Device &device, u64 slot);
void free_texture(rhi::Device &device, u64 slot);
void free_sampler(rhi::Device &device, u64 slot);
// gpu must see it. add here, drop at destroy.
void reside(rhi::Device &device, const MTL::Allocation *alloc);
void unreside(rhi::Device &device, const MTL::Allocation *alloc);
// scratch push. cmd frees it.
MTL::Buffer *push_bump_alloc(rhi::CmdBuffer &cmd, const void *data, u32 size);
} // namespace mt
namespace rhi {
struct ShaderCompiler {
Slang::ComPtr<slang::IGlobalSession> global;
Slang::ComPtr<slang::ISession> session;
bool recreate_session();
};
struct ShaderModule {
void *handle = nullptr;
rhi::ShaderStage stage = rhi::ShaderStage::NONE;
};
struct Memory {
void *handle = nullptr;
};
struct Queue {
MTL4::CommandQueue *handle = nullptr;
u32 family_index = 0;
u32 queue_index = 0;
};
struct Device {
MTL::Device *handle = nullptr;
Queue graphics;
rhi::ImageFormat format = rhi::ImageFormat::BGRA8_UNORM;
u32 min_ubo_alignment = 16;
u64 timestamp_period = 1;
mt::Heap *heap = nullptr;
};
struct Swapchain {
CA::MetalLayer *layer = nullptr;
CA::MetalDrawable *current = nullptr;
rhi::ImageFormat format = rhi::ImageFormat::BGRA8_UNORM;
rhi::Extent2D extent = {};
u32 maximumDrawableCount = 3;
bool recreate = false;
};
struct Buffer {
rhi::BufferDesc desc{};
MTL::Buffer *handle = nullptr;
void *mapped = nullptr;
u64 address = 0;
rhi::ResourceState state = rhi::ResourceState::Idle;
};
struct Image {
rhi::ImageDesc desc{};
MTL::Texture *handle = nullptr;
rhi::ResourceState state = rhi::ResourceState::Idle;
std::vector<std::pair<ImageViewDesc, MTL::Texture *>> view_cache;
};
struct ImageView {
rhi::ImageViewDesc desc{};
MTL::Texture *handle = nullptr;
u64 slot = UINT64_MAX;
};
struct BufferView {
Buffer *buffer = nullptr;
rhi::BufferViewDesc desc{};
MTL::Buffer *handle = nullptr;
u64 slot = UINT64_MAX;
};
struct Sampler {
MTL::SamplerState *handle = nullptr;
u64 slot = UINT64_MAX;
};
struct RenderPipeline {
Device *device = nullptr;
MTL::RenderPipelineState *handle = nullptr;
MTL::DepthStencilState *depth_stencil = nullptr;
// push lookup by size.
struct ConstBlock {
u32 byte_size = 0;
u32 table_index = 0;
};
std::vector<ConstBlock> const_blocks;
};
// no compute yet.
struct ComputePipeline {
Device *device = nullptr;
void *handle = nullptr;
};
struct QueryPool {
void *handle = nullptr;
rhi::QueryPoolDesc desc{};
};
struct AccelStruct {
void *handle = nullptr;
Buffer *buffer = nullptr;
};
struct ComputePipelineDesc {
Device *device = nullptr;
std::string shader_path;
std::string cs_entry;
std::vector<ShaderModule> shader_modules;
ComputePipelineDesc() = default;
ComputePipelineDesc(const ComputePipelineDesc &) = default;
ComputePipelineDesc &operator=(const ComputePipelineDesc &) = default;
ComputePipelineDesc &set_shader(ShaderCompiler &sc, const char *path, const char *cs);
ComputePipeline build();
};
struct RenderPipelineDesc {
Device *device = nullptr;
std::vector<rhi::ImageFormat> color_attachment_formats;
rhi::ImageFormat depth_attachment_format = rhi::ImageFormat::UNDEFINED;
std::string shader_path;
std::string vs_entry, fs_entry;
std::vector<ShaderModule> shader_modules;
std::string msl_source;
std::vector<VertexAttribute> vertex_attributes;
std::vector<VertexBinding> vertex_bindings;
rhi::PipelineTopology topology = rhi::PipelineTopology::TRIANGLES;
rhi::PipelineFillMode fill_mode = rhi::PipelineFillMode::SOLID;
rhi::PipelineCullMode cull_mode = rhi::PipelineCullMode::BACK;
rhi::PipelineTriangleWindingOrder winding = rhi::PipelineTriangleWindingOrder::CCW;
rhi::SampleCount samples = rhi::SampleCount::Sample1;
bool depth_test_enable = true;
bool depth_write_enable = true;
bool depth_bias_enable = false;
f32 depth_bias_constant = 0.0;
f32 depth_bias_clamp = 0.0;
f32 depth_bias_slope = 0.0;
bool depth_clamp_enable = false;
rhi::PipelineCompareOp depth_compare_op = rhi::PipelineCompareOp::LESS;
bool enable_dynamic_polygon_mode = false;
bool enable_dynamic_depth_write = false;
// reflection opt-in.
bool enable_abi_blocks = false;
bool blend_enable = false;
rhi::PipelineBlendFactor src_color_blend_factor = rhi::PipelineBlendFactor::SRC_ALPHA;
rhi::PipelineBlendFactor dst_color_blend_factor = rhi::PipelineBlendFactor::ONE_MINUS_SRC_ALPHA;
rhi::PipelineBlendOp color_blend_op = rhi::PipelineBlendOp::ADD;
rhi::PipelineBlendFactor src_alpha_blend_factor = rhi::PipelineBlendFactor::SRC_ALPHA;
rhi::PipelineBlendFactor dst_alpha_blend_factor = rhi::PipelineBlendFactor::SRC_ALPHA;
rhi::PipelineBlendOp alpha_blend_op = rhi::PipelineBlendOp::ADD;
RenderPipelineDesc() = default;
RenderPipelineDesc(const RenderPipelineDesc &) = default;
RenderPipelineDesc &operator=(const RenderPipelineDesc &) = default;
RenderPipelineDesc &set_shaders(ShaderCompiler &sc, const char *path, const char *vs, const char *fs);
RenderPipelineDesc &set_input_topology(rhi::PipelineTopology topology);
RenderPipelineDesc &add_vertex_binding(u32 binding, u32 stride = 0, bool is_instanced = false);
RenderPipelineDesc &add_vertex_attribute(u32 location, u32 binding, ImageFormat format, u32 offset = 0);
RenderPipelineDesc &set_polygon_mode(rhi::PipelineFillMode mode);
RenderPipelineDesc &set_color_format(rhi::ImageFormat format);
RenderPipelineDesc &set_depth_format(rhi::ImageFormat format);
RenderPipelineDesc &set_cull_mode(rhi::PipelineCullMode cull_mode, rhi::PipelineTriangleWindingOrder winding);
RenderPipelineDesc &set_multisampling(rhi::SampleCount count = rhi::SampleCount::Sample1);
RenderPipelineDesc &set_depth_testing(bool do_test, bool do_write, rhi::PipelineCompareOp op);
RenderPipelineDesc &set_depth_bias(bool enabled, f32 constant, f32 clamp, f32 slope);
RenderPipelineDesc &set_depth_clamp(bool enabled);
RenderPipelineDesc &set_dynamic_polygon_mode(bool enable = true) {
enable_dynamic_polygon_mode = enable;
return *this;
}
RenderPipelineDesc &set_dynamic_depth_write(bool enable = true) {
enable_dynamic_depth_write = enable;
return *this;
}
RenderPipelineDesc &set_blending(
rhi::PipelineBlendFactor src_color = rhi::PipelineBlendFactor::SRC_ALPHA,
rhi::PipelineBlendFactor dst_color = rhi::PipelineBlendFactor::ONE_MINUS_SRC_ALPHA,
rhi::PipelineBlendOp color_op = rhi::PipelineBlendOp::ADD,
rhi::PipelineBlendFactor src_alpha = rhi::PipelineBlendFactor::SRC_ALPHA,
rhi::PipelineBlendFactor dst_alpha = rhi::PipelineBlendFactor::ONE_MINUS_SRC_ALPHA,
rhi::PipelineBlendOp alpha_op = rhi::PipelineBlendOp::ADD
);
RenderPipeline build();
};
struct CmdPool;
struct CmdPool {
void *handle = nullptr;
};
struct CmdBuffer {
MTL4::CommandBuffer *handle = nullptr;
MTL4::RenderCommandEncoder *encoder = nullptr;
MTL4::CommandQueue *queue = nullptr;
MTL4::CommandAllocator *allocator = nullptr;
MTL4::ArgumentTable *table = nullptr;
MTL::ResidencySet *residency = nullptr;
bool is_rendering = false;
// scratch. freed on begin.
std::vector<MTL::Buffer *> transients;
// read at draw.
MTL::Buffer *index_buffer = nullptr;
rhi::IndexType index_type = rhi::IndexType::UINT32;
u64 index_offset = 0;
};
struct Sync {
void *timeline = nullptr;
void *device = nullptr;
};
struct SyncPoint {
Sync *sync = nullptr;
u64 value = 0;
};
struct QueueWait {
SyncPoint point;
rhi::PipelineStages stage_mask;
};
struct QueueSignal {
SyncPoint point;
};
struct QueueSubmitDesc {
u32 wait_count = 0;
const QueueWait *waits = nullptr;
u32 cmd_count = 0;
const CmdBuffer *cmds = nullptr;
u32 signal_count = 0;
const QueueSignal *signals = nullptr;
};
} // namespace rhi