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() }