shader.cpp

cross platform rendering playground

src/backend/vulkan/shader.cpp

5.6 KB
#include <volk.h>

#include <array>
#include <cstdint>
#include <cstdio>
#include <filesystem>
#include <string>

#include <cassert>

#include "backend/vulkan/debug.h"
#include "backend/vulkan/utils.h"
#include "backend/vulkan/vk_api.h"
#include "backend/vulkan/vk_conversion.h"
#include "core/globals.h"
#include "core/logger.h"

namespace rhi {
using namespace vk;

u64 get_file_last_write_time(const char *path) {
    std::error_code ec;
    auto ftime = std::filesystem::last_write_time(path, ec);
    if (ec) {
        return {};
    }
    return ftime.time_since_epoch().count();
}

static bool build_session(ShaderCompiler &sc, slang::ISession **out) {

    constexpr slang::CompilerOptionEntry emit_spirv_directly = {
        .name = slang::CompilerOptionName::EmitSpirvDirectly,
        .value{
            .kind = slang::CompilerOptionValueKind::Int,
            .intValue0 = 1,
        }
    };

    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[] = {
        emit_spirv_directly,
        optimization_level,
        debug_info_level,
    };

    slang::TargetDesc target = {};
    target.format = SLANG_SPIRV;
    target.profile = sc.global->findProfile("spirv_1_6");
    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;

#ifndef NDEBUG
    slang::PreprocessorMacroDesc debug_flag = {"DEBUG", "1"};
    desc.preprocessorMacros = &debug_flag;
    desc.preprocessorMacroCount = 1;
#endif

    return SLANG_SUCCEEDED(sc.global->createSession(desc, out));
} // namespace rhi

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...");
    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;
}

ShaderModule
load_shader(ShaderCompiler &sc, Device &device, const char *path, const char *entry_name, ShaderStage stage) {
    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(entry_name, entry_point.writeRef()))) {
        if (diagnostics) {
            VEL_ERROR(
                "Slang Find Entry Point Error ({}): {}", entry_name, (const char *)diagnostics->getBufferPointer()
            );
        } else {
            VEL_ERROR("Slang Find Entry Point Error ({}): Entry point not found", entry_name);
        }
        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> spirv;
    if (SLANG_FAILED(linked->getEntryPointCode(0, 0, spirv.writeRef(), diagnostics.writeRef()))) {
        if (diagnostics) {
            (void)std::fprintf(stderr, "%s\n", (const char *)diagnostics->getBufferPointer());
        }
        return {};
    }

    VkShaderModuleCreateInfo create_info{VK_STRUCTURE_TYPE_SHADER_MODULE_CREATE_INFO};
    create_info.codeSize = spirv->getBufferSize();
    create_info.pCode = (const u32 *)spirv->getBufferPointer();

    VkShaderModule shader = VK_NULL_HANDLE;
    VK_CHECK(vkCreateShaderModule(device.logical, &create_info, nullptr, &shader));
    {
        char mod_name[320];
        (void)snprintf(mod_name, sizeof(mod_name), "shader/%s:%s", path, entry_name);
        VK_NAME_EX(device, shader, mod_name);
    }

    VEL_INFO("Shader loaded: {} (Entry: {})", path, entry_name);
    return {shader, stage};
}

} // namespace rhi