shader.cpp

cross platform rendering playground

src/backend/metal/shader.cpp

8.28 KB
#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 = &target;
    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