Files
Dmitry Sergeev 3b51bea58a First imgui trys
2026-08-25 13:02:54 +03:00

302 lines
8.9 KiB
Go

package imgui
import (
"fmt"
"unsafe"
cimgui "github.com/AllenDang/cimgui-go/imgui"
"github.com/christerso/vulkan-go/vk"
)
// Renders ImGui draw data inside an existing Vulkan render pass
type VulkanBackend struct {
pd vk.PhysicalDevice
dev vk.Device
queue vk.Queue
cmdPool vk.CommandPool
pipeline vk.Pipeline
pipelineLayout vk.PipelineLayout
descSetLayout vk.DescriptorSetLayout
descPool vk.DescriptorPool
descSet vk.DescriptorSet
vertMod vk.ShaderModule
fragMod vk.ShaderModule
fontImg vk.AllocImage
fontView vk.ImageView
fontSampler vk.Sampler
vertBuf vk.AllocBuffer
idxBuf vk.AllocBuffer
vertSize vk.DeviceSize
idxSize vk.DeviceSize
}
// NewVulkanBackend creates the ImGui pipeline, font atlas, and buffers.
// rp is the existing render pass; format is the swapchain color format.
func NewVulkanBackend(pd vk.PhysicalDevice, dev vk.Device, queue vk.Queue, cmdPool vk.CommandPool, rp vk.RenderPass) (*VulkanBackend, error) {
b := &VulkanBackend{pd: pd, dev: dev, queue: queue, cmdPool: cmdPool}
// 1. Font atlas — build and upload via CreateTexture2D (handles staging).
io := cimgui.CurrentIO()
fontAtlas := io.Fonts()
cimgui.InternalImFontAtlasBuildMain(fontAtlas)
texData := fontAtlas.TexData()
w, h := texData.Width(), texData.Height()
pixelCount := int(w) * int(h) * 4 // RGBA32 = 4 bytes/pixel
pixels := unsafe.Slice((*byte)(unsafe.Pointer(texData.Pixels())), pixelCount)
var err error
b.fontImg, b.fontView, err = dev.CreateTexture2D(pd, queue, cmdPool, uint32(w), uint32(h), pixels)
if err != nil {
return nil, fmt.Errorf("imgui font atlas: %w", err)
}
// Font sampler: linear, clamp to edge.
b.fontSampler, err = dev.CreateSampler(vk.SamplerConfig{
MagFilter: vk.FilterLinear,
MinFilter: vk.FilterLinear,
AddressModeU: vk.SamplerAddressModeClampToEdge,
AddressModeV: vk.SamplerAddressModeClampToEdge,
})
if err != nil {
b.destroyAll()
return nil, err
}
// 2. Shaders.
b.vertMod, err = dev.CreateShaderModule(vertSPV)
if err != nil {
b.destroyAll()
return nil, err
}
b.fragMod, err = dev.CreateShaderModule(fragSPV)
if err != nil {
b.destroyAll()
return nil, err
}
// 3. Descriptor set: one combined image sampler (font), fragment stage.
b.descSetLayout, err = dev.CreateDescriptorSetLayout([]vk.DescriptorBinding{
{Binding: 0, Type: vk.DescriptorCombinedImageSampler, Count: 1, Stages: vk.ShaderStageFragment},
})
if err != nil {
b.destroyAll()
return nil, err
}
b.descPool, err = dev.CreateDescriptorPool(1, map[vk.DescriptorType]uint32{
vk.DescriptorCombinedImageSampler: 1,
})
if err != nil {
b.destroyAll()
return nil, err
}
b.descSet, err = dev.AllocateDescriptorSet(b.descPool, b.descSetLayout)
if err != nil {
b.destroyAll()
return nil, err
}
dev.UpdateImageDescriptor(b.descSet, 0, b.fontView, b.fontSampler)
// 4. Pipeline layout: one set + 16 bytes push constants (vertex+fragment).
b.pipelineLayout, err = dev.CreatePipelineLayout(
[]vk.DescriptorSetLayout{b.descSetLayout},
vk.ShaderStageVertex|vk.ShaderStageFragment, 16,
)
if err != nil {
b.destroyAll()
return nil, err
}
// 5. Graphics pipeline: vertex input (pos+uv+col = 20 bytes), alpha blend.
b.pipeline, err = dev.CreateGraphicsPipeline(vk.GraphicsPipelineConfig{
Layout: b.pipelineLayout,
RenderPass: rp,
VertexShader: b.vertMod,
FragShader: b.fragMod,
Bindings: []vk.VertexInputBinding{
{Binding: 0, Stride: 20, InputRate: vk.VertexInputRateVertex},
},
Attributes: []vk.VertexInputAttribute{
{Location: 0, Binding: 0, Format: vk.Format(103), Offset: 0}, // pos: R32G32Sfloat
{Location: 1, Binding: 0, Format: vk.Format(103), Offset: 8}, // uv: R32G32Sfloat
{Location: 2, Binding: 0, Format: vk.FormatR8G8B8A8Unorm, Offset: 16}, // col: R8G8B8A8Unorm
},
Topology: vk.TopologyTriangleList,
PolygonMode: vk.PolygonFill,
CullMode: vk.CullNone,
FrontFace: vk.FrontFaceCounterClockwise,
Blend: true,
})
if err != nil {
b.destroyAll()
return nil, err
}
// 6. Vertex/index buffers: host-visible, mapped. Start at 64KB/16KB, grow if needed.
b.vertSize = 1 << 16
b.idxSize = 1 << 14
b.vertBuf, err = dev.CreateBuffer(pd, vk.BufferConfig{
Size: b.vertSize,
Usage: vk.BufferUsageVertexBuffer,
Properties: vk.MemoryHostVisible | vk.MemoryHostCoherent,
Map: true,
})
if err != nil {
b.destroyAll()
return nil, err
}
b.idxBuf, err = dev.CreateBuffer(pd, vk.BufferConfig{
Size: vk.DeviceSize(b.idxSize),
Usage: vk.BufferUsageIndexBuffer,
Properties: vk.MemoryHostVisible | vk.MemoryHostCoherent,
Map: true,
})
if err != nil {
b.destroyAll()
return nil, err
}
return b, nil
}
// RecordDraw records ImGui draw commands into the command buffer.
// Must be called inside the render pass, after your scene draw, before EndRenderPass.
func (b *VulkanBackend) RecordDraw(cmd vk.CommandBuffer, drawData *cimgui.DrawData) {
if !drawData.Valid() || drawData.CmdListsCount() == 0 {
return
}
totalVtx := int(drawData.TotalVtxCount())
totalIdx := int(drawData.TotalIdxCount())
vertBytes := totalVtx * 20
idxBytes := totalIdx * 2
// Grow vertex buffer if needed.
if vk.DeviceSize(vertBytes) > b.vertSize {
b.dev.DestroyBuffer(b.vertBuf)
b.vertSize = vk.DeviceSize(vertBytes) * 2
b.vertBuf, _ = b.dev.CreateBuffer(b.pd, vk.BufferConfig{
Size: b.vertSize,
Usage: vk.BufferUsageVertexBuffer,
Properties: vk.MemoryHostVisible | vk.MemoryHostCoherent,
Map: true,
})
}
// Grow index buffer if needed.
if vk.DeviceSize(idxBytes) > vk.DeviceSize(b.idxBuf.Size) {
b.dev.DestroyBuffer(b.idxBuf)
b.idxBuf, _ = b.dev.CreateBuffer(b.pd, vk.BufferConfig{
Size: vk.DeviceSize(idxBytes) * 2,
Usage: vk.BufferUsageIndexBuffer,
Properties: vk.MemoryHostVisible | vk.MemoryHostCoherent,
Map: true,
})
}
// Copy all vertex/index data into the mapped buffers.
vertOffset := 0
idxOffset := 0
cmdLists := drawData.CommandLists()
for _, list := range cmdLists {
// Vertices: GetVertexBuffer returns raw C pointer + byte size.
vtxPtr, vtxBytes2 := list.GetVertexBuffer()
if vtxBytes2 > 0 {
src := unsafe.Slice((*byte)(vtxPtr), vtxBytes2)
dst := unsafe.Slice((*byte)(b.vertBuf.Mapped), vertBytes)
copy(dst[vertOffset:], src)
vertOffset += vtxBytes2
}
// Indices: GetIndexBuffer returns raw C pointer + byte size.
idxPtr, idxBytes2 := list.GetIndexBuffer()
if idxBytes2 > 0 {
src := unsafe.Slice((*byte)(idxPtr), idxBytes2)
dst := unsafe.Slice((*byte)(b.idxBuf.Mapped), idxBytes)
copy(dst[idxOffset:], src)
idxOffset += idxBytes2
}
}
// Push constants: scale + translate (transforms ImGui pixels to clip space).
displaySize := drawData.DisplaySize()
displayPos := drawData.DisplayPos()
scale := [2]float32{2.0 / displaySize.X, 2.0 / displaySize.Y}
translate := [2]float32{
-1.0 - 2.0*displayPos.X/displaySize.X,
-1.0 - 2.0*displayPos.Y/displaySize.Y,
}
pc := struct {
Scale [2]float32
Translate [2]float32
}{
Scale: scale,
Translate: translate,
}
cmd.PushConstants(b.pipelineLayout, vk.ShaderStageVertex|vk.ShaderStageFragment, 0, unsafe.Pointer(&pc), 16)
// Bind pipeline + descriptor set + vertex/index buffers.
cmd.BindPipeline(b.pipeline)
cmd.BindDescriptorSet(b.pipelineLayout, 0, b.descSet)
offsets := []vk.DeviceSize{0}
cmd.BindVertexBuffers(0, []vk.Buffer{b.vertBuf.Buffer}, offsets)
cmd.BindIndexBuffer(b.idxBuf.Buffer, 0, vk.IndexTypeUint16)
// Draw each command list, translating clip rects to scissors.
vtxOff := uint32(0)
idxOff := uint32(0)
for _, list := range cmdLists {
cmds := list.Commands()
for _, dc := range cmds {
if dc.HasUserCallback() {
dc.CallUserCallback(list)
continue
}
clip := dc.ClipRect()
sx := int32(clip.X - displayPos.X)
sy := int32(clip.Y - displayPos.Y)
ex := int32(clip.Z - displayPos.X)
ey := int32(clip.W - displayPos.Y)
if sx < 0 {
sx = 0
}
if sy < 0 {
sy = 0
}
if ex > int32(displaySize.X) {
ex = int32(displaySize.X)
}
if ey > int32(displaySize.Y) {
ey = int32(displaySize.Y)
}
if ex > sx && ey > sy {
cmd.SetScissor(vk.Rect2D{
Offset: vk.Offset2D{X: sx, Y: sy},
Extent: vk.Extent2D{Width: uint32(ex - sx), Height: uint32(ey - sy)},
})
}
cmd.DrawIndexed(dc.ElemCount(), 1, idxOff+dc.IdxOffset(), int32(vtxOff+dc.VtxOffset()), 0)
}
vtxOff += uint32(list.VtxBuffer().Size)
idxOff += uint32(list.IdxBuffer().Size)
}
}
func (b *VulkanBackend) destroyAll() {
b.dev.DestroyBuffer(b.idxBuf)
b.dev.DestroyBuffer(b.vertBuf)
b.dev.DestroyPipeline(b.pipeline)
b.dev.DestroyPipelineLayout(b.pipelineLayout)
b.dev.DestroyDescriptorPool(b.descPool)
b.dev.DestroyDescriptorSetLayout(b.descSetLayout)
b.dev.DestroyShaderModule(b.fragMod)
b.dev.DestroyShaderModule(b.vertMod)
b.dev.DestroySampler(b.fontSampler)
b.dev.DestroyImageView(b.fontView)
b.dev.DestroyImage(b.fontImg)
}
func (b *VulkanBackend) Destroy() {
b.destroyAll()
}