// // MetalScatterView.swift // Metrika // // Author: Simon-Pierre Boucher // Contact: contact@spboucher.ai // Copyright © 2026 Simon-Pierre Boucher. All rights reserved. // import MetalKit import SwiftUI import ZQGraphics /// Metal point-sprite renderer for scatter plots past the Swift Charts /// threshold (CLAUDE.md §7 pane 4). Unified memory: the packed vertex /// buffer is written once; pan (scroll) and zoom (pinch or scroll+⌥) /// only touch a 24-byte uniform. struct MetalScatterPane: View { let plot: ZQPlotSpec var body: some View { VStack(spacing: 0) { MetalScatterView(plot: plot) Divider() HStack { Text("\(plot.series.reduce(0) { $0 + $1.x.count }) points — Metal renderer") Spacer() Text("scroll to pan · pinch or ⌥-scroll to zoom · double-click to reset") } .font(.caption) .foregroundStyle(.secondary) .padding(6) } } } struct MetalScatterView: NSViewRepresentable { let plot: ZQPlotSpec func makeCoordinator() -> Renderer { Renderer() } func makeNSView(context: Context) -> ScatterMTKView { let view = ScatterMTKView() view.device = MTLCreateSystemDefaultDevice() view.delegate = context.coordinator view.renderer = context.coordinator view.enableSetNeedsDisplay = true view.isPaused = true view.clearColor = MTLClearColor(red: 0, green: 0, blue: 0, alpha: 0) view.layer?.isOpaque = false context.coordinator.attach(view: view) context.coordinator.load(plot: plot) return view } func updateNSView(_ view: ScatterMTKView, context: Context) { context.coordinator.load(plot: plot) view.needsDisplay = true } /// MTKView subclass owning the pan/zoom gestures. final class ScatterMTKView: MTKView { weak var renderer: Renderer? override func scrollWheel(with event: NSEvent) { if event.modifierFlags.contains(.option) { renderer?.zoom(by: 1 + event.scrollingDeltaY * 0.01) } else { renderer?.pan( deltaX: event.scrollingDeltaX, deltaY: event.scrollingDeltaY, viewSize: bounds.size ) } needsDisplay = true } override func magnify(with event: NSEvent) { renderer?.zoom(by: 1 + event.magnification) needsDisplay = true } override func mouseDown(with event: NSEvent) { if event.clickCount == 2 { renderer?.resetViewport() needsDisplay = true } } } /// Pipeline + vertex buffer + viewport state. @MainActor final class Renderer: NSObject, MTKViewDelegate { private struct Uniforms { var center: SIMD2 var scale: SIMD2 var pointSize: Float } private var pipeline: MTLRenderPipelineState? private var commandQueue: MTLCommandQueue? private var vertexBuffer: MTLBuffer? private var pointCount = 0 private var loadedPlotSignature = 0 // Data bounds and viewport state. private var dataCenter = SIMD2(0, 0) private var dataHalfSpan = SIMD2(1, 1) private var zoomLevel: Float = 1 private var panOffset = SIMD2(0, 0) private var aspect: Float = 1 func attach(view: MTKView) { guard let device = view.device else { return } commandQueue = device.makeCommandQueue() guard let library = device.makeDefaultLibrary(), let vertexFunction = library.makeFunction(name: "scatterVertex"), let fragmentFunction = library.makeFunction(name: "scatterFragment") else { return } let descriptor = MTLRenderPipelineDescriptor() descriptor.vertexFunction = vertexFunction descriptor.fragmentFunction = fragmentFunction descriptor.colorAttachments[0].pixelFormat = view.colorPixelFormat descriptor.colorAttachments[0].isBlendingEnabled = true descriptor.colorAttachments[0].sourceRGBBlendFactor = .sourceAlpha descriptor.colorAttachments[0].destinationRGBBlendFactor = .oneMinusSourceAlpha descriptor.colorAttachments[0].sourceAlphaBlendFactor = .one descriptor.colorAttachments[0].destinationAlphaBlendFactor = .oneMinusSourceAlpha pipeline = try? device.makeRenderPipelineState(descriptor: descriptor) } /// Rebuilds the vertex buffer when the plot actually changes. func load(plot: ZQPlotSpec) { let signature = plot.series.reduce(plot.series.count) { $0 &* 31 &+ $1.x.count } guard signature != loadedPlotSignature, let device = commandQueue?.device else { return } loadedPlotSignature = signature let palette: [SIMD4] = [ .init(79, 140, 255, 220), .init(255, 149, 0, 220), .init(52, 199, 89, 220), .init(175, 82, 222, 220), .init(255, 59, 48, 220), .init(90, 200, 250, 220), ] var minX = Float.greatestFiniteMagnitude, maxX = -Float.greatestFiniteMagnitude var minY = Float.greatestFiniteMagnitude, maxY = -Float.greatestFiniteMagnitude var packed = [UInt8]() var count = 0 for (seriesIndex, series) in plot.series.enumerated() { let color = palette[seriesIndex % palette.count] for i in 0.. 0 else { return } pointCount = count vertexBuffer = packed.withUnsafeBytes { device.makeBuffer(bytes: $0.baseAddress!, length: $0.count) } dataCenter = SIMD2((minX + maxX) / 2, (minY + maxY) / 2) dataHalfSpan = SIMD2( max((maxX - minX) / 2, 1e-9) * 1.05, max((maxY - minY) / 2, 1e-9) * 1.05 ) resetViewport() } func resetViewport() { zoomLevel = 1 panOffset = .zero } func zoom(by factor: Double) { zoomLevel = min(max(zoomLevel * Float(factor), 0.1), 10_000) } func pan(deltaX: Double, deltaY: Double, viewSize: CGSize) { guard viewSize.width > 0, viewSize.height > 0 else { return } panOffset.x -= Float(deltaX) * 2 * dataHalfSpan.x / (zoomLevel * Float(viewSize.width)) panOffset.y += Float(deltaY) * 2 * dataHalfSpan.y / (zoomLevel * Float(viewSize.height)) } nonisolated func mtkView(_ view: MTKView, drawableSizeWillChange size: CGSize) {} nonisolated func draw(in view: MTKView) { MainActor.assumeIsolated { render(in: view) } } private func render(in view: MTKView) { guard let pipeline, let vertexBuffer, pointCount > 0, let descriptor = view.currentRenderPassDescriptor, let drawable = view.currentDrawable, let commandBuffer = commandQueue?.makeCommandBuffer(), let encoder = commandBuffer.makeRenderCommandEncoder(descriptor: descriptor) else { return } var uniforms = Uniforms( center: dataCenter + panOffset, scale: SIMD2( zoomLevel / dataHalfSpan.x, zoomLevel / dataHalfSpan.y ), pointSize: Float(max(2, 6 - log10(Double(max(pointCount, 10))))) ) encoder.setRenderPipelineState(pipeline) encoder.setVertexBuffer(vertexBuffer, offset: 0, index: 0) encoder.setVertexBytes(&uniforms, length: MemoryLayout.size, index: 1) encoder.drawPrimitives(type: .point, vertexStart: 0, vertexCount: pointCount) encoder.endEncoding() commandBuffer.present(drawable) commandBuffer.commit() } } }