#include <array>
#include <cstdio>
#include <string>
#include <vector>
#include "backend/metal/mt_api.h"
#include "core/logger.h"
namespace {
// handles come with the lowering. no patching.
static bool is_msl_space(char c) {
return c == ' ' || c == '\t' || c == '\r' || c == '\n';
}
static bool is_msl_ident(char c) {
return (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '_';
}
// slang moves push around. pin it to 0.
static std::string normalize_push_index(const std::string &msl) {
std::string out = msl;
size_t i = 0;
while (i < out.size()) {
size_t p = out.find("constant* ", i);
if (p == std::string::npos) {
break;
}
size_t s = p + sizeof("constant* ") - 1;
while (s < out.size() && out[s] == ' ') {
++s;
}
if (s >= out.size() || !is_msl_ident(out[s]) || (out[s] >= '0' && out[s] <= '9')) {
i = s;
continue;
}
while (s < out.size() && is_msl_ident(out[s])) {
++s;
}
size_t t = s;
while (t < out.size() && out[t] == ' ') {
++t;
}
constexpr char kBuf[] = "[[buffer(";
if (out.compare(t, sizeof(kBuf) - 1, kBuf) != 0) {
i = s;
continue;
}
size_t d0 = t + sizeof(kBuf) - 1, d1 = d0;
while (d1 < out.size() && out[d1] >= '0' && out[d1] <= '9') {
++d1;
}
if (d1 == d0 || d1 + 3 > out.size() || out.compare(d1, 3, ")]]") != 0) {
i = s;
continue;
}
if (out.compare(d0, d1 - d0, "0") != 0) {
VEL_INFO("metal shader: push [[buffer({})]] -> [[buffer(0)]]", out.substr(d0, d1 - d0));
}
out.replace(d0, d1 - d0, "0");
i = d0 + 1;
}
return out;
}
} // namespace
namespace rhi {
static bool build_session(ShaderCompiler &sc, slang::ISession **out) {
constexpr slang::CompilerOptionEntry optimization_level = {
.name = slang::CompilerOptionName::Optimization,
.value = {
.kind = slang::CompilerOptionValueKind::Int,
.intValue0 = SLANG_OPTIMIZATION_LEVEL_NONE,
}
};
constexpr slang::CompilerOptionEntry debug_info_level = {
.name = slang::CompilerOptionName::DebugInformation,
.value = {
.kind = slang::CompilerOptionValueKind::Int,
.intValue0 = SLANG_DEBUG_INFO_LEVEL_STANDARD,
}
};
constexpr slang::CompilerOptionEntry entries[] = {
optimization_level,
debug_info_level,
};
slang::TargetDesc target = {};
target.format = SLANG_METAL;
target.flags = SLANG_TARGET_FLAG_GENERATE_WHOLE_PROGRAM;
target.compilerOptionEntries = entries;
target.compilerOptionEntryCount = std::size(entries);
const char *search_path[] = {"src/shaders"};
slang::SessionDesc desc = {};
desc.targets = ⌖
desc.targetCount = 1;
desc.allowGLSLSyntax = true;
desc.searchPaths = search_path;
desc.searchPathCount = 1;
return SLANG_SUCCEEDED(sc.global->createSession(desc, out));
}
bool ShaderCompiler::recreate_session() {
Slang::ComPtr<slang::ISession> new_session;
if (!build_session(*this, new_session.writeRef())) {
VEL_ERROR("Failed to create fresh Slang session for reload");
return false;
}
session = new_session;
return true;
}
bool init_compiler(ShaderCompiler &sc) {
VEL_INFO("Initializing Slang Shader Compiler (Metal target)...");
if (SLANG_FAILED(slang::createGlobalSession(sc.global.writeRef()))) {
VEL_ERROR("Failed to create Slang global session");
return false;
}
if (!build_session(sc, sc.session.writeRef())) {
VEL_ERROR("Failed to create Slang session");
return false;
}
return true;
}
} // namespace rhi
namespace mt {
std::string compile_msl(rhi::ShaderCompiler &sc, const char *path, const char *entryName) {
Slang::ComPtr<slang::IBlob> diagnostics;
slang::IModule *module_ptr = sc.session->loadModule(path, diagnostics.writeRef());
if (!module_ptr) {
if (diagnostics) {
VEL_ERROR("Slang Load Module Error ({}): {}", path, (const char *)diagnostics->getBufferPointer());
} else {
VEL_ERROR("Slang Load Module Error ({}): Unknown error", path);
}
return {};
}
Slang::ComPtr<slang::IModule> module(module_ptr);
Slang::ComPtr<slang::IEntryPoint> entry_point;
if (SLANG_FAILED(module->findEntryPointByName(entryName, entry_point.writeRef()))) {
if (diagnostics) {
VEL_ERROR(
"Slang Find Entry Point Error ({}): {}", entryName, (const char *)diagnostics->getBufferPointer()
);
} else {
VEL_ERROR("Slang Find Entry Point Error ({}): Entry point not found", entryName);
}
return {};
}
slang::IComponentType *components[] = {module.get(), entry_point.get()};
Slang::ComPtr<slang::IComponentType> composed;
if (SLANG_FAILED(
sc.session->createCompositeComponentType(components, 2, composed.writeRef(), diagnostics.writeRef())
)) {
if (diagnostics) {
(void)std::fprintf(stderr, "%s\n", (const char *)diagnostics->getBufferPointer());
}
return {};
}
Slang::ComPtr<slang::IComponentType> linked;
if (SLANG_FAILED(composed->link(linked.writeRef(), diagnostics.writeRef()))) {
if (diagnostics) {
(void)std::fprintf(stderr, "%s\n", (const char *)diagnostics->getBufferPointer());
}
return {};
}
Slang::ComPtr<slang::IBlob> code;
if (SLANG_FAILED(linked->getEntryPointCode(0, 0, code.writeRef(), diagnostics.writeRef()))) {
if (diagnostics) {
(void)std::fprintf(stderr, "%s\n", (const char *)diagnostics->getBufferPointer());
}
return {};
}
std::string msl(static_cast<const char *>(code->getBufferPointer()), code->getBufferSize());
return normalize_push_index(msl);
}
std::string compile_msl_linked(rhi::ShaderCompiler &sc, const char *path, const char *const *entries, u32 entry_count) {
Slang::ComPtr<slang::IBlob> diagnostics;
slang::IModule *module_ptr = sc.session->loadModule(path, diagnostics.writeRef());
if (!module_ptr) {
if (diagnostics) {
VEL_ERROR("Slang Load Module Error ({}): {}", path, (const char *)diagnostics->getBufferPointer());
} else {
VEL_ERROR("Slang Load Module Error ({}): Unknown error", path);
}
return {};
}
Slang::ComPtr<slang::IModule> module(module_ptr);
std::vector<Slang::ComPtr<slang::IEntryPoint>> ep_ptrs;
ep_ptrs.reserve(entry_count);
for (u32 i = 0; i < entry_count; ++i) {
Slang::ComPtr<slang::IEntryPoint> ep;
if (SLANG_FAILED(module->findEntryPointByName(entries[i], ep.writeRef()))) {
VEL_ERROR("Slang Find Entry Point Error ({}): {}", entries[i], path);
return {};
}
ep_ptrs.push_back(ep);
}
std::vector<slang::IComponentType *> components;
components.reserve(1 + entry_count);
components.push_back(module.get());
for (u32 i = 0; i < entry_count; ++i) {
components.push_back(ep_ptrs[i].get());
}
Slang::ComPtr<slang::IComponentType> composed;
if (SLANG_FAILED(sc.session->createCompositeComponentType(
components.data(), (u32)components.size(), composed.writeRef(), diagnostics.writeRef()
))) {
if (diagnostics) {
(void)std::fprintf(stderr, "%s\n", (const char *)diagnostics->getBufferPointer());
}
return {};
}
Slang::ComPtr<slang::IComponentType> linked;
if (SLANG_FAILED(composed->link(linked.writeRef(), diagnostics.writeRef()))) {
if (diagnostics) {
(void)std::fprintf(stderr, "%s\n", (const char *)diagnostics->getBufferPointer());
}
return {};
}
Slang::ComPtr<slang::IBlob> code;
if (SLANG_FAILED(linked->getTargetCode(0, code.writeRef(), diagnostics.writeRef()))) {
if (diagnostics) {
(void)std::fprintf(stderr, "%s\n", (const char *)diagnostics->getBufferPointer());
}
return {};
}
std::string msl(static_cast<const char *>(code->getBufferPointer()), code->getBufferSize());
return normalize_push_index(msl);
}
} // namespace mt