imgui_backend.cpp

cross platform rendering playground

src/backend/imgui_backend.cpp

9.07 KB
#include "backend/imgui_backend.h"

#include <algorithm>
#include <cstring>

#include <imgui.h>
#include <imgui_impl_glfw.h>

#include "GLFW/glfw3.h"
#include "asset/texture.h"
#include "core/logger.h"
#include "core/math_rtm.h"
#include "passes/fg.h"
#include "quick_submit.h"
#include "shaders/renderer_types.h"

namespace imgui_backend {

namespace {

rhi::Device *DEVICE = nullptr;

rhi::Image FONT_IMG = {};
rhi::ImageView FONT_VIEW = {};
u64 FONT_HANDLE = 0;

rhi::Sampler IMG_SAMP = {};
u64 IMG_SAMP_HANDLE = 0;

rhi::Buffer VTX_BUF[MAX_FRAMES_IN_FLIGHT] = {};
rhi::Buffer IDX_BUF[MAX_FRAMES_IN_FLIGHT] = {};

static void draw_callback_reset_render_state(const ImDrawList *, const ImDrawCmd *) {
}

void ensure_buffer(rhi::Buffer &buf, u64 bytes, rhi::BufferUsage usage, const char *name) {
    if (rhi::valid(buf) && buf.desc.size >= bytes) {
        return;
    }
    if (rhi::valid(buf)) {
        rhi::destroy_buffer(*DEVICE, buf);
        buf = {};
    }
    rhi::BufferDesc desc{};
    desc.size = bytes;
    desc.usage = usage;
    desc.memory = rhi::BufferMemType::Upload;
    desc.name = name;
    rhi::create_buffer(*DEVICE, desc, buf);
}

} // namespace

void init(rhi::Device &device, void *native_window, QuickSubmit &qs, Arena &staging) {
    DEVICE = &device;

    ImGui_ImplGlfw_InitForVulkan((GLFWwindow *)native_window, true);
    ImGui::GetPlatformIO().DrawCallback_ResetRenderState = draw_callback_reset_render_state;

    ImGuiIO &io = ImGui::GetIO();
    io.BackendFlags |= ImGuiBackendFlags_RendererHasVtxOffset;
    io.BackendRendererName = "gfx-playground-native";

    // font atlas upload (init-time; fonts are static).
    unsigned char *pixels = nullptr;
    int atlas_w = 0, atlas_h = 0;
    io.Fonts->GetTexDataAsRGBA32(&pixels, &atlas_w, &atlas_h);
    if (!create_texture_from_raw_rgba(
            device,
            qs,
            staging,
            FONT_IMG,
            FONT_VIEW,
            pixels,
            (u32)atlas_w,
            (u32)atlas_h,
            4,
            rhi::ImageFormat::RGBA8_UNORM,
            "font/atlas"
        )) {
        VEL_CRITICAL("imgui_backend: font atlas upload failed");
        return;
    }
    FONT_HANDLE = rhi::handle_id(device, FONT_VIEW);
    io.Fonts->SetTexID((ImTextureID)FONT_HANDLE);

    // Own static sampler; handle cached once, carried per-draw in push.
    {
        rhi::SamplerDesc sdesc{};
        sdesc.filter = rhi::SamplerFilter::LINEAR;
        sdesc.address = rhi::SamplerAddressMode::CLAMP;
        sdesc.name = "samp/imgui";
        rhi::create_sampler(device, sdesc, IMG_SAMP);
        IMG_SAMP_HANDLE = rhi::handle_id(device, IMG_SAMP);
    }

    for (u32 i = 0; i < MAX_FRAMES_IN_FLIGHT; ++i) {
        ensure_buffer(
            VTX_BUF[i], 64 * 1024, rhi::BufferUsage::VertexBuffer | rhi::BufferUsage::CopyDst, "ImGui Vertex Buffer"
        );
        ensure_buffer(
            IDX_BUF[i], 32 * 1024, rhi::BufferUsage::IndexBuffer | rhi::BufferUsage::CopyDst, "ImGui Index Buffer"
        );
    }
}

void shutdown(rhi::Device &device) {
    ImGui_ImplGlfw_Shutdown();
    for (u32 i = 0; i < MAX_FRAMES_IN_FLIGHT; ++i) {
        if (rhi::valid(VTX_BUF[i])) {
            rhi::destroy_buffer(device, VTX_BUF[i]);
            VTX_BUF[i] = {};
        }
        if (rhi::valid(IDX_BUF[i])) {
            rhi::destroy_buffer(device, IDX_BUF[i]);
            IDX_BUF[i] = {};
        }
    }
    if (rhi::valid(FONT_VIEW)) {
        rhi::destroy_image_view(device, FONT_VIEW);
        FONT_VIEW = {};
    }
    if (rhi::valid(FONT_IMG)) {
        rhi::destroy_image(device, FONT_IMG);
        FONT_IMG = {};
    }
    if (rhi::valid(IMG_SAMP)) {
        rhi::destroy_sampler(device, IMG_SAMP);
        IMG_SAMP = {};
        IMG_SAMP_HANDLE = 0;
    }
    DEVICE = nullptr;
    FONT_HANDLE = 0;
}

void begin_frame() {
    ImGui_ImplGlfw_NewFrame();
    ImGui::NewFrame();
}

u64 texture(rhi::ImageView &view) {
    return rhi::handle_id(*DEVICE, view);
}

bool prepare(const PassContext &ctx) {
    ImGui::Render();
    ImDrawData *draw_data = ImGui::GetDrawData();
    if (draw_data == nullptr || draw_data->CmdListsCount == 0 || draw_data->TotalVtxCount == 0) {
        return false;
    }

    if (draw_data->DisplaySize.x <= 0.0f || draw_data->DisplaySize.y <= 0.0f) {
        return false;
    }

    rhi::Buffer &vtx_buf = VTX_BUF[ctx.frame_index];
    rhi::Buffer &idx_buf = IDX_BUF[ctx.frame_index];
    ensure_buffer(
        vtx_buf,
        (u64)draw_data->TotalVtxCount * sizeof(ImDrawVert),
        rhi::BufferUsage::VertexBuffer | rhi::BufferUsage::CopyDst,
        "ImGui Vertex Buffer"
    );
    ensure_buffer(
        idx_buf,
        (u64)draw_data->TotalIdxCount * sizeof(ImDrawIdx),
        rhi::BufferUsage::IndexBuffer | rhi::BufferUsage::CopyDst,
        "ImGui Index Buffer"
    );

    u8 *vtx_dst = static_cast<u8 *>(vtx_buf.mapped);
    u8 *idx_dst = static_cast<u8 *>(idx_buf.mapped);
    for (int n = 0; n < draw_data->CmdListsCount; ++n) {
        const ImDrawList *list = draw_data->CmdLists[n];
        memcpy(vtx_dst, list->VtxBuffer.Data, (usize)list->VtxBuffer.Size * sizeof(ImDrawVert));
        vtx_dst += (usize)list->VtxBuffer.Size * sizeof(ImDrawVert);
        memcpy(idx_dst, list->IdxBuffer.Data, (usize)list->IdxBuffer.Size * sizeof(ImDrawIdx));
        idx_dst += (usize)list->IdxBuffer.Size * sizeof(ImDrawIdx);
    }

    rhi::barrier(ctx.cmd, vtx_buf, rhi::ResourceState::HostRead, rhi::ResourceState::VertexFetch);
    rhi::barrier(ctx.cmd, idx_buf, rhi::ResourceState::HostRead, rhi::ResourceState::IndexFetch);
    return true;
}

void draw(const PassContext &ctx, rhi::RenderPipeline &pipeline, u32 width, u32 height) {
    ImDrawData *draw_data = ImGui::GetDrawData();
    if (draw_data == nullptr || draw_data->CmdListsCount == 0 || draw_data->TotalVtxCount == 0) {
        return;
    }

    rhi::Buffer &vtx_buf = VTX_BUF[ctx.frame_index];
    rhi::Buffer &idx_buf = IDX_BUF[ctx.frame_index];

    rhi::set_pipeline(ctx.cmd, pipeline);
    rhi::set_vertex_buffer(ctx.cmd, vtx_buf, 0, 0);
    rhi::set_index_buffer(ctx.cmd, idx_buf, 0, rhi::IndexType::UINT16);

    rhi::Viewport vp{};
    vp.width = (f32)width;
    vp.height = (f32)height;
    rhi::set_viewports(ctx.cmd, &vp);

    const vec2 scale = {draw_data->FramebufferScale.x, draw_data->FramebufferScale.y};
    const vec2 translate = {
        -draw_data->DisplayPos.x * draw_data->FramebufferScale.x,
        -draw_data->DisplayPos.y * draw_data->FramebufferScale.y,
    };
    const vec2 display_size = {
        draw_data->DisplaySize.x * draw_data->FramebufferScale.x,
        draw_data->DisplaySize.y * draw_data->FramebufferScale.y,
    };

    u32 vtx_offset = 0;
    u32 idx_offset = 0;
    for (int n = 0; n < draw_data->CmdListsCount; ++n) {
        const ImDrawList *list = draw_data->CmdLists[n];
        for (int ci = 0; ci < list->CmdBuffer.Size; ++ci) {
            const ImDrawCmd *cmd = &list->CmdBuffer[ci];
            if (cmd->UserCallback != nullptr) {
                if (cmd->UserCallback == ImGui::GetPlatformIO().DrawCallback_ResetRenderState) {
                    rhi::set_pipeline(ctx.cmd, pipeline);
                    rhi::set_vertex_buffer(ctx.cmd, vtx_buf, 0, 0);
                    rhi::set_index_buffer(ctx.cmd, idx_buf, 0, rhi::IndexType::UINT16);
                    rhi::set_viewports(ctx.cmd, &vp);
                }
                continue;
            }
            if (cmd->ElemCount == 0) {
                continue;
            }
            vec4 clip = {
                (cmd->ClipRect.x - draw_data->DisplayPos.x) * draw_data->FramebufferScale.x,
                (cmd->ClipRect.y - draw_data->DisplayPos.y) * draw_data->FramebufferScale.y,
                (cmd->ClipRect.z - draw_data->DisplayPos.x) * draw_data->FramebufferScale.x,
                (cmd->ClipRect.w - draw_data->DisplayPos.y) * draw_data->FramebufferScale.y,
            };
            if (clip.x >= clip.z || clip.y >= clip.w) {
                continue;
            }
            rhi::Rect scissor{};
            scissor.x = (i16)std::max(clip.x, 0.0f);
            scissor.y = (i16)std::max(clip.y, 0.0f);
            scissor.width = (i16)std::min(clip.z, display_size.x) - scissor.x;
            scissor.height = (i16)std::min(clip.w, display_size.y) - scissor.y;
            if (scissor.width <= 0 || scissor.height <= 0) {
                continue;
            }
            rhi::set_scissors(ctx.cmd, &scissor);

            ImGuiPushConstants pc{};
            pc.scale = scale;
            pc.translate = translate;
            pc.display_size = display_size;
            pc.tex = (u32)(u64)cmd->GetTexID();
            pc.samp = IMG_SAMP_HANDLE;
            {
                static_assert(sizeof(pc) <= 128, "push constants exceed global 128B range");
                char pc128[128] = {};
                std::memcpy(pc128, &pc, sizeof(pc));
                rhi::set_constants(ctx.cmd, pipeline, rhi::ShaderStage::ALL, sizeof(pc128), pc128);
            }
            rhi::draw_indexed(ctx.cmd, cmd->ElemCount, 1, idx_offset + cmd->IdxOffset, vtx_offset + cmd->VtxOffset, 0);
        }
        vtx_offset += (u32)list->VtxBuffer.Size;
        idx_offset += (u32)list->IdxBuffer.Size;
    }
}

} // namespace imgui_backend