#include <math.h>
#ifndef INFINITY
#define INFINITY __builtin_huge_valf()
#endif
#include <Foundation/Foundation.hpp>
#include <Metal/Metal.hpp>
#include "backend/metal/mt_api.h"
#include "core/logger.h"
namespace {
// has texture/sampler?
static bool struct_has_opaque_members(MTL::StructType *st) {
if (!st || !st->members()) {
return false;
}
auto *members = st->members();
for (NS::UInteger i = 0; i < members->count(); ++i) {
MTL::StructMember *m = static_cast<MTL::StructMember *>(members->object(i));
MTL::DataType dt = m->dataType();
if (dt == MTL::DataTypeTexture || dt == MTL::DataTypeSampler) {
return true;
}
if (dt == MTL::DataTypeStruct) {
if (struct_has_opaque_members(m->structType())) {
return true;
}
}
if (dt == MTL::DataTypeArray) {
MTL::ArrayType *at = m->arrayType();
if (!at) {
continue;
}
MTL::DataType et = at->elementType();
if (et == MTL::DataTypeTexture || et == MTL::DataTypeSampler) {
return true;
}
if (et == MTL::DataTypeStruct) {
if (struct_has_opaque_members(at->elementStructType())) {
return true;
}
}
}
}
return false;
}
static bool arg_struct_has_opaque_members(MTL::Argument *arg) {
return struct_has_opaque_members(arg->bufferStructType());
}
template <typename TArray> static void record_stage_args(rhi::RenderPipeline &out, TArray *args, const char *stage) {
for (NS::UInteger i = 0; i < args->count(); ++i) {
MTL::Argument *a = static_cast<MTL::Argument *>(args->object(i));
if (a->type() != MTL::ArgumentTypeBuffer) {
continue;
}
MTL::StructType *st = a->bufferStructType();
if (!st) {
continue;
}
if (arg_struct_has_opaque_members(a)) {
continue;
}
rhi::RenderPipeline::ConstBlock cb{};
cb.byte_size = (u32)a->bufferDataSize();
cb.table_index = (u32)a->index();
bool known = false;
for (auto &e : out.const_blocks) {
if (e.table_index == cb.table_index) {
known = true;
break;
}
}
if (!known) {
out.const_blocks.push_back(cb);
VEL_INFO("metal const: {} bytes at table index {} ({} stage)", cb.byte_size, cb.table_index, stage);
}
}
}
} // namespace
namespace rhi {
static NS::String *str(const std::string &s) {
return NS::String::string(s.c_str(), NS::UTF8StringEncoding);
}
static MTL::VertexFormat to_vertex_format(rhi::ImageFormat format) {
switch (format) {
case rhi::ImageFormat::R32_FLOAT:
return MTL::VertexFormatFloat;
case rhi::ImageFormat::RG32_FLOAT:
return MTL::VertexFormatFloat2;
case rhi::ImageFormat::RGB32_FLOAT:
return MTL::VertexFormatFloat3;
case rhi::ImageFormat::RGBA32_FLOAT:
return MTL::VertexFormatFloat4;
default:
return MTL::VertexFormatFloat2;
}
}
static MTL::PixelFormat to_pixel_format(rhi::ImageFormat format) {
switch (format) {
case rhi::ImageFormat::RGBA8_UNORM:
return MTL::PixelFormatRGBA8Unorm;
case rhi::ImageFormat::RGBA8_SRGB:
return MTL::PixelFormatRGBA8Unorm_sRGB;
case rhi::ImageFormat::BGRA8_UNORM:
return MTL::PixelFormatBGRA8Unorm;
default:
return MTL::PixelFormatBGRA8Unorm;
}
}
static MTL::CullMode to_cull_mode(rhi::PipelineCullMode mode) {
switch (mode) {
case rhi::PipelineCullMode::NONE:
return MTL::CullModeNone;
case rhi::PipelineCullMode::FRONT:
return MTL::CullModeFront;
case rhi::PipelineCullMode::BACK:
default:
return MTL::CullModeBack;
}
}
static MTL::Winding to_winding(rhi::PipelineTriangleWindingOrder winding) {
switch (winding) {
case rhi::PipelineTriangleWindingOrder::CW:
return MTL::WindingClockwise;
case rhi::PipelineTriangleWindingOrder::CCW:
default:
return MTL::WindingCounterClockwise;
}
}
static MTL::PixelFormat to_depth_format(rhi::ImageFormat format) {
switch (format) {
case rhi::ImageFormat::D32_FLOAT:
return MTL::PixelFormatDepth32Float;
default:
return MTL::PixelFormatInvalid;
}
}
static MTL::CompareFunction to_compare(rhi::PipelineCompareOp op) {
switch (op) {
case rhi::PipelineCompareOp::NEVER:
return MTL::CompareFunctionNever;
case rhi::PipelineCompareOp::EQUAL:
return MTL::CompareFunctionEqual;
case rhi::PipelineCompareOp::LESS_EQUAL:
return MTL::CompareFunctionLessEqual;
case rhi::PipelineCompareOp::GREATER:
return MTL::CompareFunctionGreater;
case rhi::PipelineCompareOp::NOT_EQUAL:
return MTL::CompareFunctionNotEqual;
case rhi::PipelineCompareOp::GREATER_EQUAL:
return MTL::CompareFunctionGreaterEqual;
case rhi::PipelineCompareOp::ALWAYS:
return MTL::CompareFunctionAlways;
case rhi::PipelineCompareOp::LESS:
default:
return MTL::CompareFunctionLess;
}
}
static rhi::ImageFormat first_color_format(const RenderPipelineDesc &desc) {
if (!desc.color_attachment_formats.empty()) {
return desc.color_attachment_formats[0];
}
return rhi::ImageFormat::BGRA8_UNORM;
}
RenderPipelineDesc &
RenderPipelineDesc::set_shaders(ShaderCompiler &sc, const char *path, const char *vs, const char *fs) {
shader_path = path ? path : "";
vs_entry = vs ? vs : "";
fs_entry = fs ? fs : "";
const char *entries[] = {vs, fs};
msl_source = mt::compile_msl_linked(sc, path, entries, 2);
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_input_topology(rhi::PipelineTopology topo) {
topology = topo;
return *this;
}
RenderPipelineDesc &
RenderPipelineDesc::add_vertex_attribute(u32 location, u32 binding, rhi::ImageFormat format, u32 offset) {
vertex_attributes.push_back({location, binding, format, offset});
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::add_vertex_binding(u32 binding, u32 stride, bool is_instanced) {
if (binding >= mt::kMaxVertexBindings) {
VEL_ERROR("add_vertex_binding: binding {} exceeds max {}", binding, mt::kMaxVertexBindings);
return *this;
}
vertex_bindings.push_back({binding, stride, is_instanced});
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_polygon_mode(rhi::PipelineFillMode mode) {
fill_mode = mode;
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_color_format(rhi::ImageFormat format) {
color_attachment_formats.clear();
color_attachment_formats.push_back(format);
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_depth_format(rhi::ImageFormat format) {
depth_attachment_format = format;
return *this;
}
RenderPipelineDesc &
RenderPipelineDesc::set_cull_mode(rhi::PipelineCullMode mode, rhi::PipelineTriangleWindingOrder wind) {
cull_mode = mode;
winding = wind;
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_multisampling(rhi::SampleCount count) {
samples = count;
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_depth_testing(bool do_test, bool do_write, rhi::PipelineCompareOp op) {
depth_test_enable = do_test;
depth_write_enable = do_write;
depth_compare_op = op;
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_depth_bias(bool enabled, f32 constant, f32 clamp, f32 slope) {
depth_bias_enable = enabled;
depth_bias_constant = constant;
depth_bias_clamp = clamp;
depth_bias_slope = slope;
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_depth_clamp(bool enabled) {
depth_clamp_enable = enabled;
return *this;
}
RenderPipelineDesc &RenderPipelineDesc::set_blending(
rhi::PipelineBlendFactor src_color,
rhi::PipelineBlendFactor dst_color,
rhi::PipelineBlendOp color_op,
rhi::PipelineBlendFactor src_alpha,
rhi::PipelineBlendFactor dst_alpha,
rhi::PipelineBlendOp alpha_op
) {
blend_enable = true;
src_color_blend_factor = src_color;
dst_color_blend_factor = dst_color;
color_blend_op = color_op;
src_alpha_blend_factor = src_alpha;
dst_alpha_blend_factor = dst_alpha;
alpha_blend_op = alpha_op;
return *this;
}
ComputePipelineDesc &ComputePipelineDesc::set_shader(ShaderCompiler &sc, const char *path, const char *cs) {
(void)sc;
shader_path = path ? path : "";
cs_entry = cs ? cs : "";
VEL_ERROR("ComputePipeline is not implemented on Metal");
return *this;
}
ComputePipeline ComputePipelineDesc::build() {
VEL_ERROR("ComputePipeline::build is not implemented on Metal");
ComputePipeline out{};
out.device = device;
return out;
}
RenderPipeline RenderPipelineDesc::build() {
RenderPipeline out{};
out.device = device;
MTL::Device *dev = device ? device->handle : nullptr;
if (!dev) {
VEL_ERROR("RenderPipeline::build: null device");
return out;
}
NS::Error *error = nullptr;
MTL::Library *lib = dev->newLibrary(str(msl_source), nullptr, &error);
if (!lib) {
VEL_ERROR(
"Failed to compile Metal library: {}", error ? error->localizedDescription()->utf8String() : "unknown error"
);
return out;
}
MTL::Function *vs = lib->newFunction(str(vs_entry));
MTL::Function *fs = lib->newFunction(str(fs_entry));
if (!vs || !fs) {
VEL_ERROR("Failed to find entry points '{}'/'{}' in Metal library", vs_entry, fs_entry);
lib->release();
return out;
}
MTL::RenderPipelineDescriptor *rpd = MTL::RenderPipelineDescriptor::alloc()->init();
rpd->setVertexFunction(vs);
rpd->setFragmentFunction(fs);
rpd->colorAttachments()->object(0)->setPixelFormat(to_pixel_format(first_color_format(*this)));
// pipeline depth must match pass.
if (depth_attachment_format != rhi::ImageFormat::UNDEFINED) {
rpd->setDepthAttachmentPixelFormat(to_depth_format(depth_attachment_format));
}
if (!vertex_attributes.empty()) {
// fetch at 2+binding. matches cmd.
MTL::VertexDescriptor *vd = MTL::VertexDescriptor::vertexDescriptor();
for (const auto &attr : vertex_attributes) {
MTL::VertexAttributeDescriptor *a = vd->attributes()->object(attr.location);
a->setFormat(to_vertex_format(attr.format));
a->setOffset(attr.offset);
a->setBufferIndex(mt::kVertexBaseIndex + attr.binding);
}
for (const auto &binding : vertex_bindings) {
MTL::VertexBufferLayoutDescriptor *l = vd->layouts()->object(mt::kVertexBaseIndex + binding.binding);
l->setStride(binding.stride);
l->setStepFunction(
binding.is_instanced ? MTL::VertexStepFunctionPerInstance : MTL::VertexStepFunctionPerVertex
);
}
rpd->setVertexDescriptor(vd);
}
(void)to_cull_mode(cull_mode);
(void)to_winding(winding);
MTL::RenderPipelineState *ps = nullptr;
if (enable_abi_blocks) {
MTL::RenderPipelineReflection *refl = nullptr;
ps = dev->newRenderPipelineState(rpd, MTL::PipelineOptionArgumentInfo, &refl, &error);
if (!ps || !refl) {
VEL_ERROR(
"Failed to create render pipeline state: {}",
error ? error->localizedDescription()->utf8String() : "unknown error"
);
vs->release();
fs->release();
rpd->release();
lib->release();
return out;
}
record_stage_args(out, refl->vertexArguments(), "vertex");
record_stage_args(out, refl->fragmentArguments(), "fragment");
refl->release();
} else {
ps = dev->newRenderPipelineState(rpd, &error);
if (!ps) {
VEL_ERROR(
"Failed to create render pipeline state: {}",
error ? error->localizedDescription()->utf8String() : "unknown error"
);
vs->release();
fs->release();
rpd->release();
lib->release();
return out;
}
}
vs->release();
fs->release();
rpd->release();
lib->release();
out.handle = ps;
if (depth_attachment_format != rhi::ImageFormat::UNDEFINED && (depth_test_enable || depth_write_enable)) {
MTL::DepthStencilDescriptor *dd = MTL::DepthStencilDescriptor::alloc()->init();
dd->setDepthCompareFunction(to_compare(depth_compare_op));
dd->setDepthWriteEnabled(depth_write_enable);
out.depth_stencil = dev->newDepthStencilState(dd);
dd->release();
if (!out.depth_stencil) {
VEL_ERROR("RenderPipeline::build: newDepthStencilState failed (depth runs untested)");
}
}
return out;
}
} // namespace rhi